mathtext2doc 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.
- mathtext2doc/__init__.py +12 -0
- mathtext2doc/__main__.py +5 -0
- mathtext2doc/cli.py +249 -0
- mathtext2doc/compiler.py +495 -0
- mathtext2doc/parser.py +809 -0
- mathtext2doc/plotter.py +775 -0
- mathtext2doc/texgen.py +228 -0
- mathtext2doc-0.1.0.dist-info/METADATA +328 -0
- mathtext2doc-0.1.0.dist-info/RECORD +13 -0
- mathtext2doc-0.1.0.dist-info/WHEEL +5 -0
- mathtext2doc-0.1.0.dist-info/entry_points.txt +2 -0
- mathtext2doc-0.1.0.dist-info/licenses/LICENSE +21 -0
- mathtext2doc-0.1.0.dist-info/top_level.txt +1 -0
mathtext2doc/plotter.py
ADDED
|
@@ -0,0 +1,775 @@
|
|
|
1
|
+
"""函数图绘制器。
|
|
2
|
+
|
|
3
|
+
输入:parser.Plot(一组 PlotItem)
|
|
4
|
+
输出:PNG 文件路径
|
|
5
|
+
|
|
6
|
+
支持:
|
|
7
|
+
- 2D 显函数 y = f(x)
|
|
8
|
+
- 2D 隐函数 F(x, y) = 0
|
|
9
|
+
- 多函数同图(显隐混合)
|
|
10
|
+
- 自动分配颜色
|
|
11
|
+
- 标签默认标注在曲线可见部分的几何中点偏上
|
|
12
|
+
- 标签之间自动避让(基于 bbox 重叠检测的迭代位移)
|
|
13
|
+
- 中文 label(matplotlib 使用 Noto Sans SC)
|
|
14
|
+
- 常见初等函数:+ - * / ^、sin cos tan log ln exp sqrt abs pi e
|
|
15
|
+
|
|
16
|
+
实现说明:
|
|
17
|
+
- 显函数:sympy 解析表达式 → lambdify → numpy linspace 采样
|
|
18
|
+
- 隐函数:sympy 解析 F(x, y) = 0 → 移项为 F(x, y) → numpy meshgrid
|
|
19
|
+
+ matplotlib contour(level=0) 绘制零等高线
|
|
20
|
+
- 标签避让:先按曲线可见中点初始化每个标签位置;然后用迭代算法
|
|
21
|
+
检查所有标签 bbox 是否重叠,重叠则沿曲线方向滑动或上移,最多迭代 N 次。
|
|
22
|
+
"""
|
|
23
|
+
|
|
24
|
+
from __future__ import annotations
|
|
25
|
+
|
|
26
|
+
import math
|
|
27
|
+
import os
|
|
28
|
+
from dataclasses import dataclass
|
|
29
|
+
from typing import List, Optional, Tuple
|
|
30
|
+
|
|
31
|
+
import matplotlib
|
|
32
|
+
matplotlib.use("Agg") # 无显示设备
|
|
33
|
+
import matplotlib.font_manager as fm
|
|
34
|
+
|
|
35
|
+
# 注册中文字体(自动探测常见路径)
|
|
36
|
+
_FONT_CANDIDATES = [
|
|
37
|
+
# Linux 常见路径
|
|
38
|
+
"/usr/share/fonts/truetype/chinese/NotoSansSC[wght].ttf",
|
|
39
|
+
"/usr/share/fonts/truetype/chinese/NotoSansSC-Regular.ttf",
|
|
40
|
+
"/usr/share/fonts/truetype/noto/NotoSansCJK-Regular.ttc",
|
|
41
|
+
"/usr/share/fonts/opentype/noto/NotoSansCJK-Regular.ttc",
|
|
42
|
+
"/usr/share/fonts/truetype/wqy/wqy-zenhei.ttc",
|
|
43
|
+
"/usr/share/fonts/truetype/wqy/wqy-microhei.ttc",
|
|
44
|
+
"/usr/share/fonts/truetype/lxgw-wenkai/LXGWWenKai-Regular.ttf",
|
|
45
|
+
# macOS 常见路径
|
|
46
|
+
"/System/Library/Fonts/PingFang.ttc",
|
|
47
|
+
"/Library/Fonts/Songti.ttc",
|
|
48
|
+
# DejaVu Sans 兜底(拉丁 + 符号)
|
|
49
|
+
"/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf",
|
|
50
|
+
]
|
|
51
|
+
for _f in _FONT_CANDIDATES:
|
|
52
|
+
if os.path.exists(_f):
|
|
53
|
+
try:
|
|
54
|
+
fm.fontManager.addfont(_f)
|
|
55
|
+
except Exception:
|
|
56
|
+
pass
|
|
57
|
+
|
|
58
|
+
# 探测可用的中文字体名(按优先级)
|
|
59
|
+
_AVAILABLE_CN_FONTS = []
|
|
60
|
+
for _name in [
|
|
61
|
+
"Noto Sans SC", "Noto Sans CJK SC", "WenQuanYi Zen Hei",
|
|
62
|
+
"WenQuanYi Micro Hei", "LXGW WenKai", "PingFang SC", "Songti SC",
|
|
63
|
+
"Sarasa Mono SC",
|
|
64
|
+
]:
|
|
65
|
+
try:
|
|
66
|
+
if fm.findfont(_name, fallback_to_default=False) != fm.findfont("DejaVu Sans"):
|
|
67
|
+
_AVAILABLE_CN_FONTS.append(_name)
|
|
68
|
+
except Exception:
|
|
69
|
+
pass
|
|
70
|
+
|
|
71
|
+
import matplotlib.pyplot as plt
|
|
72
|
+
import numpy as np
|
|
73
|
+
import sympy as sp
|
|
74
|
+
|
|
75
|
+
# 优先中文字体,然后 DejaVu Sans 兜底
|
|
76
|
+
plt.rcParams["font.sans-serif"] = (_AVAILABLE_CN_FONTS or ["Noto Sans SC"]) + [
|
|
77
|
+
"DejaVu Sans",
|
|
78
|
+
]
|
|
79
|
+
plt.rcParams["axes.unicode_minus"] = False
|
|
80
|
+
|
|
81
|
+
from .parser import Plot, PlotItem
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
# ---------------------------------------------------------------------------
|
|
85
|
+
# 颜色循环
|
|
86
|
+
# ---------------------------------------------------------------------------
|
|
87
|
+
|
|
88
|
+
# 一组对比明显、色盲友好的颜色
|
|
89
|
+
_COLOR_CYCLE = [
|
|
90
|
+
"#1f77b4", # 蓝
|
|
91
|
+
"#d62728", # 红
|
|
92
|
+
"#2ca02c", # 绿
|
|
93
|
+
"#ff7f0e", # 橙
|
|
94
|
+
"#9467bd", # 紫
|
|
95
|
+
"#17becf", # 青
|
|
96
|
+
"#8c564b", # 棕
|
|
97
|
+
"#e377c2", # 粉
|
|
98
|
+
]
|
|
99
|
+
|
|
100
|
+
|
|
101
|
+
# ---------------------------------------------------------------------------
|
|
102
|
+
# 表达式解析
|
|
103
|
+
# ---------------------------------------------------------------------------
|
|
104
|
+
|
|
105
|
+
# sympy 默认支持 log (=ln)、sin、cos、tan、exp、sqrt、abs、pi、E。
|
|
106
|
+
# 我们额外提供 `ln` 作为 `log` 的别名,`e` 作为 E 的别名。
|
|
107
|
+
_LOCAL_NAMES = {
|
|
108
|
+
"pi": sp.pi,
|
|
109
|
+
"e": sp.E,
|
|
110
|
+
"E": sp.E,
|
|
111
|
+
"ln": sp.log,
|
|
112
|
+
"log": sp.log, # sympy 中 log 默认是自然对数
|
|
113
|
+
"log10": sp.log,
|
|
114
|
+
"sin": sp.sin,
|
|
115
|
+
"cos": sp.cos,
|
|
116
|
+
"tan": sp.tan,
|
|
117
|
+
"exp": sp.exp,
|
|
118
|
+
"sqrt": sp.sqrt,
|
|
119
|
+
"abs": sp.Abs,
|
|
120
|
+
}
|
|
121
|
+
|
|
122
|
+
|
|
123
|
+
def _to_sympy_expr(s: str) -> sp.Expr:
|
|
124
|
+
"""把用户表达式转成 sympy 表达式。把 ^ 替换为 **。"""
|
|
125
|
+
s = s.replace("^", "**")
|
|
126
|
+
try:
|
|
127
|
+
return sp.sympify(s, locals=_LOCAL_NAMES)
|
|
128
|
+
except Exception as ex:
|
|
129
|
+
raise PlotRenderError(f"无法解析表达式 {s!r}:{ex}")
|
|
130
|
+
|
|
131
|
+
|
|
132
|
+
# ---------------------------------------------------------------------------
|
|
133
|
+
# 显函数 / 隐函数绘制
|
|
134
|
+
# ---------------------------------------------------------------------------
|
|
135
|
+
|
|
136
|
+
@dataclass
|
|
137
|
+
class _PlottedCurve:
|
|
138
|
+
"""已经绘制好的曲线信息,用于标签避让。"""
|
|
139
|
+
label: Optional[str]
|
|
140
|
+
color: str
|
|
141
|
+
# 曲线上的采样点(用于找标签锚点)。numpy 数组,可能是 NaN。
|
|
142
|
+
xs: np.ndarray
|
|
143
|
+
ys: np.ndarray
|
|
144
|
+
# 标签候选锚点(在数据坐标里)
|
|
145
|
+
anchor: Optional[Tuple[float, float]] = None
|
|
146
|
+
# 最终标签锚点(经过避让后)
|
|
147
|
+
final_anchor: Optional[Tuple[float, float]] = None
|
|
148
|
+
|
|
149
|
+
|
|
150
|
+
class PlotRenderError(Exception):
|
|
151
|
+
pass
|
|
152
|
+
|
|
153
|
+
|
|
154
|
+
def _plot_explicit(ax, item: PlotItem, color: str) -> _PlottedCurve:
|
|
155
|
+
"""绘制 y = f(x)。"""
|
|
156
|
+
# 从 "y = sin(x)" 中取出右边的 f(x)
|
|
157
|
+
if "=" not in item.expr:
|
|
158
|
+
raise PlotRenderError(f"显函数缺少 '=':{item.expr!r}")
|
|
159
|
+
lhs, rhs = item.expr.split("=", 1)
|
|
160
|
+
lhs = lhs.strip()
|
|
161
|
+
rhs = rhs.strip()
|
|
162
|
+
if lhs != "y":
|
|
163
|
+
raise PlotRenderError(
|
|
164
|
+
f"显函数左边必须是 'y',实际为 {lhs!r}(在 {item.expr!r} 中)"
|
|
165
|
+
)
|
|
166
|
+
expr = _to_sympy_expr(rhs)
|
|
167
|
+
x = sp.symbols("x")
|
|
168
|
+
try:
|
|
169
|
+
f = sp.lambdify(x, expr, modules=["numpy", _LOCAL_NAMES])
|
|
170
|
+
except Exception as ex:
|
|
171
|
+
raise PlotRenderError(f"显函数 lambdify 失败 {rhs!r}:{ex}")
|
|
172
|
+
|
|
173
|
+
a, b = item.x_range
|
|
174
|
+
a = float(a)
|
|
175
|
+
b = float(b)
|
|
176
|
+
n = max(400, int((b - a) * 80))
|
|
177
|
+
n = min(n, 4000)
|
|
178
|
+
xs = np.linspace(a, b, n)
|
|
179
|
+
try:
|
|
180
|
+
ys_raw = f(xs)
|
|
181
|
+
# 常数函数 / 标量返回:广播到与 xs 同形状
|
|
182
|
+
if np.isscalar(ys_raw) or (hasattr(ys_raw, "shape") and ys_raw.shape == ()):
|
|
183
|
+
ys = np.full_like(xs, float(ys_raw))
|
|
184
|
+
else:
|
|
185
|
+
ys = np.asarray(ys_raw, dtype=float)
|
|
186
|
+
# 形状不匹配(如某些 sympy 函数返回 (1, N))→ flatten
|
|
187
|
+
if ys.shape != xs.shape:
|
|
188
|
+
ys = ys.reshape(xs.shape) if ys.size == xs.size else np.broadcast_to(ys, xs.shape).astype(float)
|
|
189
|
+
except Exception as ex:
|
|
190
|
+
raise PlotRenderError(f"显函数求值失败 {rhs!r}:{ex}")
|
|
191
|
+
|
|
192
|
+
# 屏蔽 NaN / inf
|
|
193
|
+
mask = np.isfinite(ys)
|
|
194
|
+
xs_plot = xs[mask]
|
|
195
|
+
ys_plot = ys[mask]
|
|
196
|
+
|
|
197
|
+
# 限制 ys 在合理范围内(避免 tan 等函数在大值处画穿)
|
|
198
|
+
# 我们把 |y| > 1e6 的点设为 NaN,让 matplotlib 断开线段
|
|
199
|
+
ys_limited = ys.copy()
|
|
200
|
+
ys_limited[np.abs(ys_limited) > 1e6] = np.nan
|
|
201
|
+
ax.plot(xs, ys_limited, color=color, linewidth=1.6, label=item.label or "")
|
|
202
|
+
|
|
203
|
+
# 计算可见部分几何中点
|
|
204
|
+
if len(xs_plot) > 0:
|
|
205
|
+
# 用弧长加权中点:找累计弧长 50% 处的点
|
|
206
|
+
dx = np.diff(xs_plot)
|
|
207
|
+
dy = np.diff(ys_plot)
|
|
208
|
+
seglen = np.sqrt(dx * dx + dy * dy)
|
|
209
|
+
cum = np.concatenate([[0], np.cumsum(seglen)])
|
|
210
|
+
if cum[-1] > 0:
|
|
211
|
+
half = cum[-1] / 2.0
|
|
212
|
+
idx = int(np.searchsorted(cum, half))
|
|
213
|
+
idx = min(idx, len(xs_plot) - 1)
|
|
214
|
+
anchor = (float(xs_plot[idx]), float(ys_plot[idx]))
|
|
215
|
+
else:
|
|
216
|
+
anchor = (float(xs_plot[0]), float(ys_plot[0]))
|
|
217
|
+
else:
|
|
218
|
+
anchor = ((a + b) / 2.0, 0.0)
|
|
219
|
+
|
|
220
|
+
return _PlottedCurve(
|
|
221
|
+
label=item.label, color=color, xs=xs_plot, ys=ys_plot, anchor=anchor
|
|
222
|
+
)
|
|
223
|
+
|
|
224
|
+
|
|
225
|
+
def _plot_implicit(ax, item: PlotItem, color: str) -> _PlottedCurve:
|
|
226
|
+
"""绘制 F(x, y) = 0。"""
|
|
227
|
+
if "=" not in item.expr:
|
|
228
|
+
raise PlotRenderError(f"隐函数缺少 '=':{item.expr!r}")
|
|
229
|
+
lhs, rhs = item.expr.split("=", 1)
|
|
230
|
+
lhs_expr = _to_sympy_expr(lhs.strip())
|
|
231
|
+
rhs_expr = _to_sympy_expr(rhs.strip())
|
|
232
|
+
F = lhs_expr - rhs_expr
|
|
233
|
+
x, y = sp.symbols("x y")
|
|
234
|
+
try:
|
|
235
|
+
f = sp.lambdify((x, y), F, modules=["numpy", _LOCAL_NAMES])
|
|
236
|
+
except Exception as ex:
|
|
237
|
+
raise PlotRenderError(f"隐函数 lambdify 失败 {item.expr!r}:{ex}")
|
|
238
|
+
|
|
239
|
+
a, b = item.x_range
|
|
240
|
+
c, d = item.y_range # 已在 parser 中校验非 None
|
|
241
|
+
a, b = float(a), float(b)
|
|
242
|
+
c, d = float(c), float(d)
|
|
243
|
+
|
|
244
|
+
n = 400
|
|
245
|
+
xs = np.linspace(a, b, n)
|
|
246
|
+
ys = np.linspace(c, d, n)
|
|
247
|
+
X, Y = np.meshgrid(xs, ys)
|
|
248
|
+
try:
|
|
249
|
+
Z = np.asarray(f(X, Y), dtype=float)
|
|
250
|
+
except Exception as ex:
|
|
251
|
+
raise PlotRenderError(f"隐函数求值失败 {item.expr!r}:{ex}")
|
|
252
|
+
|
|
253
|
+
# 屏蔽 NaN / inf
|
|
254
|
+
Z = np.where(np.isfinite(Z), Z, np.nan)
|
|
255
|
+
|
|
256
|
+
# 用 contour 画 0 等高线
|
|
257
|
+
try:
|
|
258
|
+
cs = ax.contour(X, Y, Z, levels=[0], colors=[color], linewidths=1.6)
|
|
259
|
+
except Exception as ex:
|
|
260
|
+
raise PlotRenderError(f"隐函数 contour 绘制失败 {item.expr!r}:{ex}")
|
|
261
|
+
|
|
262
|
+
# 取出 contour 的路径段,计算几何中点
|
|
263
|
+
# matplotlib 3.9+:ContourSet 自身就是 Collection;旧版本有 .collections
|
|
264
|
+
if hasattr(cs, "get_paths"):
|
|
265
|
+
try:
|
|
266
|
+
paths = cs.get_paths()
|
|
267
|
+
except Exception:
|
|
268
|
+
paths = []
|
|
269
|
+
elif hasattr(cs, "collections") and cs.collections:
|
|
270
|
+
try:
|
|
271
|
+
paths = cs.collections[0].get_paths()
|
|
272
|
+
except Exception:
|
|
273
|
+
paths = []
|
|
274
|
+
else:
|
|
275
|
+
paths = []
|
|
276
|
+
|
|
277
|
+
all_pts = []
|
|
278
|
+
for p in paths:
|
|
279
|
+
v = p.vertices
|
|
280
|
+
if len(v) > 0:
|
|
281
|
+
all_pts.append(v)
|
|
282
|
+
if all_pts:
|
|
283
|
+
pts = np.vstack(all_pts)
|
|
284
|
+
# 弧长中点
|
|
285
|
+
dx = np.diff(pts[:, 0])
|
|
286
|
+
dy = np.diff(pts[:, 1])
|
|
287
|
+
seglen = np.sqrt(dx * dx + dy * dy)
|
|
288
|
+
cum = np.concatenate([[0], np.cumsum(seglen)])
|
|
289
|
+
if cum[-1] > 0:
|
|
290
|
+
half = cum[-1] / 2.0
|
|
291
|
+
idx = int(np.searchsorted(cum, half))
|
|
292
|
+
idx = min(idx, len(pts) - 1)
|
|
293
|
+
anchor = (float(pts[idx, 0]), float(pts[idx, 1]))
|
|
294
|
+
else:
|
|
295
|
+
anchor = (float(pts[0, 0]), float(pts[0, 1]))
|
|
296
|
+
xs_plot = pts[:, 0]
|
|
297
|
+
ys_plot = pts[:, 1]
|
|
298
|
+
else:
|
|
299
|
+
anchor = ((a + b) / 2.0, (c + d) / 2.0)
|
|
300
|
+
xs_plot = np.array([])
|
|
301
|
+
ys_plot = np.array([])
|
|
302
|
+
|
|
303
|
+
return _PlottedCurve(
|
|
304
|
+
label=item.label, color=color, xs=xs_plot, ys=ys_plot, anchor=anchor
|
|
305
|
+
)
|
|
306
|
+
|
|
307
|
+
|
|
308
|
+
# ---------------------------------------------------------------------------
|
|
309
|
+
# 几何图形渲染
|
|
310
|
+
# ---------------------------------------------------------------------------
|
|
311
|
+
|
|
312
|
+
def _plot_shape(ax, item: PlotItem, color: str) -> _PlottedCurve:
|
|
313
|
+
"""渲染几何图形。"""
|
|
314
|
+
shape = item.shape
|
|
315
|
+
params = item.params
|
|
316
|
+
label = item.label
|
|
317
|
+
anchor = None
|
|
318
|
+
xs_plot = np.array([])
|
|
319
|
+
ys_plot = np.array([])
|
|
320
|
+
|
|
321
|
+
if shape == "point":
|
|
322
|
+
x, y = params["at"]
|
|
323
|
+
ax.plot([x], [y], marker="o", color=color, markersize=6,
|
|
324
|
+
markeredgecolor=color, markerfacecolor=color)
|
|
325
|
+
anchor = (float(x), float(y))
|
|
326
|
+
xs_plot = np.array([float(x)])
|
|
327
|
+
ys_plot = np.array([float(y)])
|
|
328
|
+
|
|
329
|
+
elif shape == "segment":
|
|
330
|
+
x1, y1 = params["from"]
|
|
331
|
+
x2, y2 = params["to"]
|
|
332
|
+
ax.plot([x1, x2], [y1, y2], color=color, linewidth=1.8,
|
|
333
|
+
solid_capstyle="round")
|
|
334
|
+
# 端点小圆点
|
|
335
|
+
ax.plot([x1, x2], [y1, y2], marker="o", color=color, markersize=4,
|
|
336
|
+
linestyle="None")
|
|
337
|
+
anchor = (float((x1 + x2) / 2), float((y1 + y2) / 2))
|
|
338
|
+
xs_plot = np.array([float(x1), float(x2)])
|
|
339
|
+
ys_plot = np.array([float(y1), float(y2)])
|
|
340
|
+
|
|
341
|
+
elif shape == "line":
|
|
342
|
+
# 直线:过两点,但要延伸到 axes 边界
|
|
343
|
+
x1, y1 = params["from"]
|
|
344
|
+
x2, y2 = params["to"]
|
|
345
|
+
# 计算方向向量,延伸到当前 xlim 的两端
|
|
346
|
+
dx, dy = x2 - x1, y2 - y1
|
|
347
|
+
if dx == 0 and dy == 0:
|
|
348
|
+
raise PlotRenderError(f"直线 from 和 to 不能重合:{item.params}")
|
|
349
|
+
# 用参数 t 延伸:t=0 在 from,t=1 在 to,延伸到 t=-100..100 兜底
|
|
350
|
+
# 实际延伸长度由 set_xlim 后的 clip 决定
|
|
351
|
+
ts = np.array([-1000, 1000])
|
|
352
|
+
xs = x1 + ts * dx
|
|
353
|
+
ys = y1 + ts * dy
|
|
354
|
+
ax.plot(xs, ys, color=color, linewidth=1.4)
|
|
355
|
+
anchor = (float((x1 + x2) / 2), float((y1 + y2) / 2))
|
|
356
|
+
xs_plot = np.array([float(x1), float(x2)])
|
|
357
|
+
ys_plot = np.array([float(y1), float(y2)])
|
|
358
|
+
|
|
359
|
+
elif shape == "circle":
|
|
360
|
+
cx, cy = params["center"]
|
|
361
|
+
r = float(params["r"])
|
|
362
|
+
if r <= 0:
|
|
363
|
+
raise PlotRenderError(f"圆半径必须为正:r={r}")
|
|
364
|
+
theta = np.linspace(0, 2 * np.pi, 200)
|
|
365
|
+
xs = cx + r * np.cos(theta)
|
|
366
|
+
ys = cy + r * np.sin(theta)
|
|
367
|
+
ax.plot(xs, ys, color=color, linewidth=1.6)
|
|
368
|
+
# 圆心小点
|
|
369
|
+
ax.plot([cx], [cy], marker="o", color=color, markersize=3)
|
|
370
|
+
anchor = (float(cx), float(cy + r)) # 顶部
|
|
371
|
+
xs_plot = xs
|
|
372
|
+
ys_plot = ys
|
|
373
|
+
|
|
374
|
+
elif shape == "ellipse":
|
|
375
|
+
cx, cy = params["center"]
|
|
376
|
+
a = float(params["a"])
|
|
377
|
+
b = float(params["b"])
|
|
378
|
+
if a <= 0 or b <= 0:
|
|
379
|
+
raise PlotRenderError(f"椭圆半轴必须为正:a={a}, b={b}")
|
|
380
|
+
theta = np.linspace(0, 2 * np.pi, 200)
|
|
381
|
+
xs = cx + a * np.cos(theta)
|
|
382
|
+
ys = cy + b * np.sin(theta)
|
|
383
|
+
ax.plot(xs, ys, color=color, linewidth=1.6)
|
|
384
|
+
ax.plot([cx], [cy], marker="o", color=color, markersize=3)
|
|
385
|
+
anchor = (float(cx), float(cy + b))
|
|
386
|
+
xs_plot = xs
|
|
387
|
+
ys_plot = ys
|
|
388
|
+
|
|
389
|
+
elif shape == "polygon":
|
|
390
|
+
pts = params["points"]
|
|
391
|
+
if len(pts) < 3:
|
|
392
|
+
raise PlotRenderError(f"多边形至少需要 3 个顶点,得到 {len(pts)}")
|
|
393
|
+
xs = [p[0] for p in pts] + [pts[0][0]]
|
|
394
|
+
ys = [p[1] for p in pts] + [pts[0][1]]
|
|
395
|
+
ax.plot(xs, ys, color=color, linewidth=1.6)
|
|
396
|
+
ax.plot([p[0] for p in pts], [p[1] for p in pts],
|
|
397
|
+
marker="o", color=color, markersize=4, linestyle="None")
|
|
398
|
+
# 几何中心
|
|
399
|
+
cx = sum(p[0] for p in pts) / len(pts)
|
|
400
|
+
cy = sum(p[1] for p in pts) / len(pts)
|
|
401
|
+
anchor = (float(cx), float(cy))
|
|
402
|
+
xs_plot = np.array([float(x) for x in xs])
|
|
403
|
+
ys_plot = np.array([float(y) for y in ys])
|
|
404
|
+
|
|
405
|
+
elif shape == "rectangle":
|
|
406
|
+
x0, y0 = params["origin"]
|
|
407
|
+
w = float(params["w"])
|
|
408
|
+
h = float(params["h"])
|
|
409
|
+
if w <= 0 or h <= 0:
|
|
410
|
+
raise PlotRenderError(f"矩形宽高必须为正:w={w}, h={h}")
|
|
411
|
+
xs = [x0, x0 + w, x0 + w, x0, x0]
|
|
412
|
+
ys = [y0, y0, y0 + h, y0 + h, y0]
|
|
413
|
+
ax.plot(xs, ys, color=color, linewidth=1.6)
|
|
414
|
+
anchor = (float(x0 + w / 2), float(y0 + h / 2))
|
|
415
|
+
xs_plot = np.array([float(x) for x in xs])
|
|
416
|
+
ys_plot = np.array([float(y) for y in ys])
|
|
417
|
+
|
|
418
|
+
elif shape == "vector":
|
|
419
|
+
x1, y1 = params["from"]
|
|
420
|
+
x2, y2 = params["to"]
|
|
421
|
+
# 用 ax.annotate 画带箭头的向量
|
|
422
|
+
ax.annotate(
|
|
423
|
+
"",
|
|
424
|
+
xy=(x2, y2), xytext=(x1, y1),
|
|
425
|
+
arrowprops=dict(arrowstyle="->", color=color, lw=1.8),
|
|
426
|
+
)
|
|
427
|
+
anchor = (float((x1 + x2) / 2), float((y1 + y2) / 2))
|
|
428
|
+
xs_plot = np.array([float(x1), float(x2)])
|
|
429
|
+
ys_plot = np.array([float(y1), float(y2)])
|
|
430
|
+
|
|
431
|
+
elif shape == "parabola":
|
|
432
|
+
# 标准方程(vertex=(h,k), p=焦距, direction=up/down/left/right):
|
|
433
|
+
# up: (x-h)^2 = 4p(y-k) → y = (x-h)^2/(4p) + k
|
|
434
|
+
# down: (x-h)^2 = -4p(y-k) → y = -(x-h)^2/(4p) + k
|
|
435
|
+
# right: (y-k)^2 = 4p(x-h) → x = (y-k)^2/(4p) + h
|
|
436
|
+
# left: (y-k)^2 = -4p(x-h) → x = -(y-k)^2/(4p) + h
|
|
437
|
+
h, k = params["vertex"]
|
|
438
|
+
p = float(params["p"])
|
|
439
|
+
direction = str(params.get("direction", "up")).lower().strip()
|
|
440
|
+
if p == 0:
|
|
441
|
+
raise PlotRenderError("parabola 的 p 不能为 0")
|
|
442
|
+
# 在 vertex 周围画 4 倍 p 的范围(够看清形状又不会太远)
|
|
443
|
+
span = abs(p) * 4 + 1.5
|
|
444
|
+
if direction in ("up", "down"):
|
|
445
|
+
sign = 1 if direction == "up" else -1
|
|
446
|
+
xs = np.linspace(float(h) - span, float(h) + span, 400)
|
|
447
|
+
ys = sign * (xs - float(h)) ** 2 / (4 * p) + float(k)
|
|
448
|
+
ax.plot(xs, ys, color=color, linewidth=1.6)
|
|
449
|
+
# 顶点
|
|
450
|
+
ax.plot([h], [k], marker="o", color=color, markersize=4)
|
|
451
|
+
anchor = (float(h), float(k) + sign * abs(p))
|
|
452
|
+
xs_plot = xs
|
|
453
|
+
ys_plot = ys
|
|
454
|
+
elif direction in ("left", "right"):
|
|
455
|
+
sign = 1 if direction == "right" else -1
|
|
456
|
+
ys = np.linspace(float(k) - span, float(k) + span, 400)
|
|
457
|
+
xs = sign * (ys - float(k)) ** 2 / (4 * p) + float(h)
|
|
458
|
+
ax.plot(xs, ys, color=color, linewidth=1.6)
|
|
459
|
+
ax.plot([h], [k], marker="o", color=color, markersize=4)
|
|
460
|
+
anchor = (float(h) + sign * abs(p), float(k))
|
|
461
|
+
xs_plot = xs
|
|
462
|
+
ys_plot = ys
|
|
463
|
+
else:
|
|
464
|
+
raise PlotRenderError(
|
|
465
|
+
f"parabola 的 direction 必须是 up/down/left/right,得到 {direction!r}"
|
|
466
|
+
)
|
|
467
|
+
|
|
468
|
+
else:
|
|
469
|
+
raise PlotRenderError(f"未知几何图形:{shape!r}")
|
|
470
|
+
|
|
471
|
+
return _PlottedCurve(
|
|
472
|
+
label=label, color=color, xs=xs_plot, ys=ys_plot, anchor=anchor
|
|
473
|
+
)
|
|
474
|
+
|
|
475
|
+
|
|
476
|
+
# ---------------------------------------------------------------------------
|
|
477
|
+
# 标签避让
|
|
478
|
+
# ---------------------------------------------------------------------------
|
|
479
|
+
|
|
480
|
+
def _resolve_label_anchor(
|
|
481
|
+
ax, curves: List[_PlottedCurve], renderer
|
|
482
|
+
) -> None:
|
|
483
|
+
"""对每条有 label 的曲线计算最终锚点。
|
|
484
|
+
|
|
485
|
+
算法:
|
|
486
|
+
1. 初始锚点 = 曲线几何中点,向上偏移 0.04 * (y_max - y_min)。
|
|
487
|
+
2. 把数据坐标转成显示坐标(pixels),用文本 bbox 检测重叠。
|
|
488
|
+
3. 重叠时,对后插入的标签,沿 8 个方向(上、右上、右、右下、下、左下、左、左上)
|
|
489
|
+
螺旋外扩,直到不重叠或达到最大迭代次数。
|
|
490
|
+
4. 若仍重叠,则把标签放在轴外上方,作为兜底。
|
|
491
|
+
"""
|
|
492
|
+
if not curves:
|
|
493
|
+
return
|
|
494
|
+
|
|
495
|
+
# 数据范围
|
|
496
|
+
all_xmin = min(np.min(c.xs) if len(c.xs) else math.inf for c in curves if c.anchor)
|
|
497
|
+
all_xmax = max(np.max(c.xs) if len(c.xs) else -math.inf for c in curves if c.anchor)
|
|
498
|
+
all_ymin = min(np.min(c.ys) if len(c.ys) else math.inf for c in curves if c.anchor)
|
|
499
|
+
all_ymax = max(np.max(c.ys) if len(c.ys) else -math.inf for c in curves if c.anchor)
|
|
500
|
+
if not all(math.isfinite(v) for v in [all_xmin, all_xmax, all_ymin, all_ymax]):
|
|
501
|
+
# 兜底
|
|
502
|
+
y_offset = 0.1
|
|
503
|
+
else:
|
|
504
|
+
y_offset = 0.04 * (all_ymax - all_ymin) if all_ymax > all_ymin else 0.1
|
|
505
|
+
|
|
506
|
+
placed_boxes = [] # 已放置标签的 bbox(display coords)
|
|
507
|
+
|
|
508
|
+
for c in curves:
|
|
509
|
+
if not c.label or c.anchor is None:
|
|
510
|
+
c.final_anchor = c.anchor
|
|
511
|
+
continue
|
|
512
|
+
ax_pt = (c.anchor[0], c.anchor[1] + y_offset)
|
|
513
|
+
# 转 display 坐标
|
|
514
|
+
disp = ax.transData.transform(ax_pt)
|
|
515
|
+
|
|
516
|
+
# 估算文本 bbox
|
|
517
|
+
# 用 renderer 测量
|
|
518
|
+
text_obj = ax.text(
|
|
519
|
+
ax_pt[0], ax_pt[1], c.label,
|
|
520
|
+
color=c.color, fontsize=10,
|
|
521
|
+
ha="center", va="bottom",
|
|
522
|
+
)
|
|
523
|
+
try:
|
|
524
|
+
bb = text_obj.get_window_extent(renderer=renderer)
|
|
525
|
+
except Exception:
|
|
526
|
+
bb = None
|
|
527
|
+
|
|
528
|
+
# 螺旋避让
|
|
529
|
+
if bb is not None:
|
|
530
|
+
best_bb = bb
|
|
531
|
+
best_disp = disp
|
|
532
|
+
best_pt = ax_pt
|
|
533
|
+
directions = [
|
|
534
|
+
(0, 1), (1, 1), (1, 0), (1, -1),
|
|
535
|
+
(0, -1), (-1, -1), (-1, 0), (-1, 1),
|
|
536
|
+
]
|
|
537
|
+
step = 8 # 像素
|
|
538
|
+
max_iter = 40
|
|
539
|
+
ok = False
|
|
540
|
+
for it in range(max_iter):
|
|
541
|
+
overlap = False
|
|
542
|
+
for prev in placed_boxes:
|
|
543
|
+
if _bbox_overlap(best_bb, prev):
|
|
544
|
+
overlap = True
|
|
545
|
+
break
|
|
546
|
+
if not overlap:
|
|
547
|
+
ok = True
|
|
548
|
+
break
|
|
549
|
+
# 沿外扩方向移动
|
|
550
|
+
d = directions[it % len(directions)]
|
|
551
|
+
radius = step * (1 + it // len(directions))
|
|
552
|
+
new_disp = (best_disp[0] + d[0] * radius, best_disp[1] + d[1] * radius)
|
|
553
|
+
# 回到数据坐标
|
|
554
|
+
new_pt = ax.transData.inverted().transform(new_disp)
|
|
555
|
+
text_obj.set_position((float(new_pt[0]), float(new_pt[1])))
|
|
556
|
+
try:
|
|
557
|
+
best_bb = text_obj.get_window_extent(renderer=renderer)
|
|
558
|
+
except Exception:
|
|
559
|
+
pass
|
|
560
|
+
best_disp = new_disp
|
|
561
|
+
best_pt = (float(new_pt[0]), float(new_pt[1]))
|
|
562
|
+
|
|
563
|
+
placed_boxes.append(best_bb)
|
|
564
|
+
c.final_anchor = best_pt
|
|
565
|
+
if not ok:
|
|
566
|
+
# 已经尽力,保留当前位置
|
|
567
|
+
pass
|
|
568
|
+
else:
|
|
569
|
+
c.final_anchor = ax_pt
|
|
570
|
+
placed_boxes.append(None)
|
|
571
|
+
|
|
572
|
+
|
|
573
|
+
def _bbox_overlap(a, b) -> bool:
|
|
574
|
+
if a is None or b is None:
|
|
575
|
+
return False
|
|
576
|
+
# matplotlib Bbox:x0,y0,x1,y1
|
|
577
|
+
return not (a.x1 < b.x0 or b.x1 < a.x0 or a.y1 < b.y0 or b.y1 < a.y0)
|
|
578
|
+
|
|
579
|
+
|
|
580
|
+
# ---------------------------------------------------------------------------
|
|
581
|
+
# 主入口
|
|
582
|
+
# ---------------------------------------------------------------------------
|
|
583
|
+
|
|
584
|
+
def render_plot(
|
|
585
|
+
plot: Plot,
|
|
586
|
+
out_path: str,
|
|
587
|
+
dpi: int = 150,
|
|
588
|
+
display_width: Optional[float] = None,
|
|
589
|
+
) -> str:
|
|
590
|
+
r"""渲染一张 plot 到 PNG。
|
|
591
|
+
|
|
592
|
+
参数:
|
|
593
|
+
plot: Plot 节点
|
|
594
|
+
out_path: 输出 PNG 路径
|
|
595
|
+
dpi: 目标有效 DPI(每英寸显示长度的像素数)。若 display_width 给定,
|
|
596
|
+
实际 savefig dpi 会自动调整以保持此有效分辨率一致。
|
|
597
|
+
display_width: 图在文档中的显示宽度(相对于 \paperwidth,0~1)。
|
|
598
|
+
给定时启用 auto-DPI;为 None 时直接用 dpi 作为 savefig dpi。
|
|
599
|
+
|
|
600
|
+
返回 out_path。
|
|
601
|
+
"""
|
|
602
|
+
if not plot.items:
|
|
603
|
+
raise PlotRenderError("@plot 指令没有任何绘制项")
|
|
604
|
+
|
|
605
|
+
# ---------- 计算 savefig dpi(auto-DPI)----------
|
|
606
|
+
# 思路:无论图显示多大,保证"每英寸显示长度的像素数"≈ dpi,
|
|
607
|
+
# 这样大图小图都有相同的视觉清晰度,不浪费像素也不糊。
|
|
608
|
+
# display_inches = display_width × paperwidth_inches
|
|
609
|
+
# pixel_width = display_inches × dpi
|
|
610
|
+
# savefig_dpi = pixel_width / figsize_width
|
|
611
|
+
_FIGSIZE_WIDTH = 5.0 # 当前 figsize=(5, 4) 的宽度
|
|
612
|
+
_PAPERWIDTH_INCHES = 8.27 # A4 纸宽 21cm
|
|
613
|
+
_MIN_DPI = 80 # 下限:避免小图文字锯齿
|
|
614
|
+
_MAX_DPI = 400 # 上限:避免大图文件过大
|
|
615
|
+
if display_width is not None:
|
|
616
|
+
display_inches = display_width * _PAPERWIDTH_INCHES
|
|
617
|
+
auto_dpi = display_inches * dpi / _FIGSIZE_WIDTH
|
|
618
|
+
savefig_dpi = max(_MIN_DPI, min(_MAX_DPI, auto_dpi))
|
|
619
|
+
else:
|
|
620
|
+
savefig_dpi = float(dpi)
|
|
621
|
+
|
|
622
|
+
# ---------- 决定坐标范围 ----------
|
|
623
|
+
# 函数曲线项必须有 x_range(parser 已校验),几何图形项的 x_range 可选
|
|
624
|
+
# 收集:用户显式指定的范围 + 几何图形的图形边界
|
|
625
|
+
user_xmin = math.inf
|
|
626
|
+
user_xmax = -math.inf
|
|
627
|
+
user_ymin = math.inf
|
|
628
|
+
user_ymax = -math.inf
|
|
629
|
+
has_user_x = False
|
|
630
|
+
has_user_y = False
|
|
631
|
+
has_implicit = False
|
|
632
|
+
has_shape = False
|
|
633
|
+
|
|
634
|
+
for it in plot.items:
|
|
635
|
+
if it.kind == "implicit":
|
|
636
|
+
has_implicit = True
|
|
637
|
+
if it.kind == "shape":
|
|
638
|
+
has_shape = True
|
|
639
|
+
if it.x_range is not None:
|
|
640
|
+
a, b = it.x_range
|
|
641
|
+
user_xmin = min(user_xmin, float(a))
|
|
642
|
+
user_xmax = max(user_xmax, float(b))
|
|
643
|
+
has_user_x = True
|
|
644
|
+
if it.y_range is not None:
|
|
645
|
+
c, d = it.y_range
|
|
646
|
+
user_ymin = min(user_ymin, float(c))
|
|
647
|
+
user_ymax = max(user_ymax, float(d))
|
|
648
|
+
has_user_y = True
|
|
649
|
+
|
|
650
|
+
fig, ax = plt.subplots(figsize=(5, 4), constrained_layout=True)
|
|
651
|
+
|
|
652
|
+
# ---------- 渲染所有项 ----------
|
|
653
|
+
curves: List[_PlottedCurve] = []
|
|
654
|
+
for i, item in enumerate(plot.items):
|
|
655
|
+
color = _COLOR_CYCLE[i % len(_COLOR_CYCLE)]
|
|
656
|
+
if item.kind == "explicit":
|
|
657
|
+
cur = _plot_explicit(ax, item, color)
|
|
658
|
+
elif item.kind == "implicit":
|
|
659
|
+
cur = _plot_implicit(ax, item, color)
|
|
660
|
+
elif item.kind == "shape":
|
|
661
|
+
cur = _plot_shape(ax, item, color)
|
|
662
|
+
else:
|
|
663
|
+
raise PlotRenderError(f"未知绘图类型:{item.kind!r}")
|
|
664
|
+
curves.append(cur)
|
|
665
|
+
|
|
666
|
+
# ---------- 设置坐标范围 ----------
|
|
667
|
+
if has_user_x and has_user_y:
|
|
668
|
+
# 用户显式指定了 x、y 范围
|
|
669
|
+
xmin, xmax = user_xmin, user_xmax
|
|
670
|
+
ymin, ymax = user_ymin, user_ymax
|
|
671
|
+
elif has_user_x and not has_user_y:
|
|
672
|
+
# 用户只指定了 x 范围,y 由 autoscale 决定
|
|
673
|
+
xmin, xmax = user_xmin, user_xmax
|
|
674
|
+
ax.relim()
|
|
675
|
+
ax.autoscale_view()
|
|
676
|
+
ymin, ymax = ax.get_ylim()
|
|
677
|
+
ax.set_xlim(xmin, xmax)
|
|
678
|
+
elif has_shape and not has_user_x:
|
|
679
|
+
# 纯几何图形,无任何范围指定:从图形数据自动估算
|
|
680
|
+
ax.relim()
|
|
681
|
+
ax.autoscale_view()
|
|
682
|
+
xmin, xmax = ax.get_xlim()
|
|
683
|
+
ymin, ymax = ax.get_ylim()
|
|
684
|
+
else:
|
|
685
|
+
# 兜底(不应该到这)
|
|
686
|
+
ax.relim()
|
|
687
|
+
ax.autoscale_view()
|
|
688
|
+
xmin, xmax = ax.get_xlim()
|
|
689
|
+
ymin, ymax = ax.get_ylim()
|
|
690
|
+
|
|
691
|
+
# 给范围留 8% 边距(避免图形贴边)
|
|
692
|
+
x_span = xmax - xmin
|
|
693
|
+
y_span = ymax - ymin
|
|
694
|
+
if x_span > 0:
|
|
695
|
+
pad = x_span * 0.08
|
|
696
|
+
xmin -= pad
|
|
697
|
+
xmax += pad
|
|
698
|
+
if y_span > 0:
|
|
699
|
+
pad = y_span * 0.08
|
|
700
|
+
ymin -= pad
|
|
701
|
+
ymax += pad
|
|
702
|
+
|
|
703
|
+
ax.set_xlim(xmin, xmax)
|
|
704
|
+
ax.set_ylim(ymin, ymax)
|
|
705
|
+
|
|
706
|
+
# ---------- 决定 aspect ----------
|
|
707
|
+
# 隐函数图、几何图形、显隐混合 → equal(几何意义正确)
|
|
708
|
+
# 纯显函数图 → equal 除非 y 跨度 > 3 倍 x 跨度(避免 tan 等压扁)
|
|
709
|
+
if has_implicit or has_shape:
|
|
710
|
+
ax.set_aspect("equal", adjustable="box")
|
|
711
|
+
else:
|
|
712
|
+
x_span_final = xmax - xmin
|
|
713
|
+
y_span_final = ymax - ymin
|
|
714
|
+
if x_span_final > 0 and y_span_final / x_span_final <= 3.0:
|
|
715
|
+
ax.set_aspect("equal", adjustable="box")
|
|
716
|
+
|
|
717
|
+
# ---------- 增强坐标系绘制 ----------
|
|
718
|
+
_draw_axes(ax, xmin, xmax, ymin, ymax)
|
|
719
|
+
|
|
720
|
+
# 标签避让(必须先 draw 一次才能拿到 renderer)
|
|
721
|
+
fig.canvas.draw()
|
|
722
|
+
renderer = fig.canvas.get_renderer()
|
|
723
|
+
_resolve_label_anchor(ax, curves, renderer)
|
|
724
|
+
|
|
725
|
+
fig.savefig(out_path, dpi=savefig_dpi)
|
|
726
|
+
plt.close(fig)
|
|
727
|
+
return out_path
|
|
728
|
+
|
|
729
|
+
|
|
730
|
+
def _draw_axes(ax, xmin, xmax, ymin, ymax) -> None:
|
|
731
|
+
"""绘制增强坐标系:箭头轴线、原点 O、x/y 轴标签。
|
|
732
|
+
|
|
733
|
+
与 matplotlib 默认 spines 不同,这里用 axhline/axvline + annotate 画带箭头的轴,
|
|
734
|
+
更接近中学/大学数学教材的坐标系画法。
|
|
735
|
+
"""
|
|
736
|
+
# 隐藏默认 spines
|
|
737
|
+
for spine in ax.spines.values():
|
|
738
|
+
spine.set_visible(False)
|
|
739
|
+
|
|
740
|
+
# 网格
|
|
741
|
+
ax.grid(True, linewidth=0.4, alpha=0.4, color="#ccc")
|
|
742
|
+
|
|
743
|
+
# 坐标轴:用浅色线穿过原点(如果原点在视野内)
|
|
744
|
+
ax.axhline(0, color="#444", linewidth=1.0, zorder=1)
|
|
745
|
+
ax.axvline(0, color="#444", linewidth=1.0, zorder=1)
|
|
746
|
+
|
|
747
|
+
# x 轴箭头(右端)
|
|
748
|
+
ax.annotate(
|
|
749
|
+
"",
|
|
750
|
+
xy=(xmax, 0), xytext=(xmax - (xmax - xmin) * 0.04, 0),
|
|
751
|
+
arrowprops=dict(arrowstyle="->", color="#444", lw=1.2),
|
|
752
|
+
annotation_clip=False,
|
|
753
|
+
)
|
|
754
|
+
# y 轴箭头(上端)
|
|
755
|
+
ax.annotate(
|
|
756
|
+
"",
|
|
757
|
+
xy=(0, ymax), xytext=(0, ymax - (ymax - ymin) * 0.04),
|
|
758
|
+
arrowprops=dict(arrowstyle="->", color="#444", lw=1.2),
|
|
759
|
+
annotation_clip=False,
|
|
760
|
+
)
|
|
761
|
+
|
|
762
|
+
# 轴标签:x 在右端下方,y 在上端左侧
|
|
763
|
+
ax.text(xmax, 0, " x", ha="left", va="bottom", fontsize=11,
|
|
764
|
+
color="#222", clip_on=False)
|
|
765
|
+
ax.text(0, ymax, "y ", ha="right", va="top", fontsize=11,
|
|
766
|
+
color="#222", clip_on=False)
|
|
767
|
+
|
|
768
|
+
# 原点 O(仅当原点在视野内且不在边缘时显示)
|
|
769
|
+
if xmin < 0 < xmax and ymin < 0 < ymax:
|
|
770
|
+
ax.text(0, 0, " O", ha="left", va="top", fontsize=9,
|
|
771
|
+
color="#222", clip_on=False)
|
|
772
|
+
|
|
773
|
+
# 刻度
|
|
774
|
+
ax.tick_params(axis="both", which="both", direction="out",
|
|
775
|
+
top=False, right=False, labelsize=8, colors="#444")
|