forked from QuantFans/quantdigger
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathwidgets.py
More file actions
690 lines (590 loc) · 24.7 KB
/
Copy pathwidgets.py
File metadata and controls
690 lines (590 loc) · 24.7 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
# -*- coding: utf-8 -*-
import sys
import six
from six.moves import range
from matplotlib.widgets import AxesWidget
from matplotlib.widgets import MultiCursor
from matplotlib.ticker import Formatter
import matplotlib.ticker as mticker
import numpy as np
from quantdigger.util.log import gen_log as log
def slider_strtime_format(delta):
""" 根据时间间隔判断周期及slider上相应的显示形式 """
if delta.days >= 1:
return '%Y-%m'
elif delta.seconds == 60:
return '%H:%M'
else:
# 日内其它分钟
return '%H:%M'
class Slider(AxesWidget):
"""
A slider representing a floating point range
The following attributes are defined
*ax* : the slider :class:`matplotlib.axes.Axes` instance
*val* : the current slider value
*vline* : a :class:`matplotlib.lines.Line2D` instance
representing the initial value of the slider
*poly* : A :class:`matplotlib.patches.Polygon` instance
which is the slider knob
*valfmt* : the format string for formatting the slider text
*label* : a :class:`matplotlib.text.Text` instance
for the slider label
*closedmin* : whether the slider is closed on the minimum
*closedmax* : whether the slider is closed on the maximum
*slidermin* : another slider - if not *None*, this slider must be
greater than *slidermin*
*slidermax* : another slider - if not *None*, this slider must be
less than *slidermax*
*drag_enabled* : allow for mouse dragging on slider
Call :meth:`add_observer` to connect to the slider event
"""
def __init__(self, ax, name, label, valmin, valmax, valinit=0.5, width=1, valfmt='%1.2f',
time_index = None, closedmin=True, closedmax=True, slidermin=None,
slidermax=None, drag_enabled=True, **kwargs):
"""
Create a slider from *valmin* to *valmax* in axes *ax*
*valinit*
The slider initial position
*label*
The slider label
*valfmt*
Used to format the slider value
*closedmin* and *closedmax*
Indicate whether the slider interval is closed
*slidermin* and *slidermax*
Used to constrain the value of this slider to the values
of other sliders.
additional kwargs are passed on to ``self.poly`` which is the
:class:`matplotlib.patches.Rectangle` which draws the slider
knob. See the :class:`matplotlib.patches.Rectangle` documentation
valid property names (e.g., *facecolor*, *edgecolor*, *alpha*, ...)
"""
AxesWidget.__init__(self, ax)
self.label = ax.text(-0.02, 0.5, label, transform=ax.transAxes,
verticalalignment='center',
horizontalalignment='right')
self.valtext = None
self.poly = None
self.reinit(valmin, valmax, valinit, width, valfmt, time_index, **kwargs)
self.name = name
self.cnt = 0
self.closedmin = closedmin
self.closedmax = closedmax
self.slidermin = slidermin
self.slidermax = slidermax
self.drag_active = False
self.drag_enabled = drag_enabled
self.observers = {}
ax.set_yticks([])
#ax.set_xticks([]) # disable ticks
ax.set_navigate(False)
self._connect()
def _xticks_to_display(self, valmax):
interval = valmax / 5
v = 0
xticks = []
for i in range(0, 6):
xticks.append(v)
v += interval
return xticks
def _value_format(self, x):
"""docstring for timess"""
ind = int(round(x))
if ind>=len(self._index) or ind<0: return ''
return self._index[ind].strftime(self._fmt)
self._slider = Slider(0, self._data_length-1,
self._data_length-1, self._data_length/50, "%d",
self._data.index)
def reinit(self, valmin, valmax, valinit=0.5, width=1, valfmt='%1.2f',
time_index = None, **kwargs):
""" [valmin, valmax] """
self.ax.set_xticks(self._xticks_to_display(valmax))
self._index = time_index
self.valmin = valmin
self.valmax = valmax
self.val = valinit
self.valinit = valinit
self.width = width
self.valfmt = valfmt
self._fmt = slider_strtime_format(time_index[1] - time_index[0])
self.ax.set_xlim((valmin, valmax))
self._data_length = valmax
if self.valtext:
self.valtext.remove()
if self.poly:
self.poly.remove()
# 滑动条的形状
self.poly = self.ax.axvspan(valmax-self.width/2,valmax+self.width/2, 0, 1, **kwargs)
#axhspan
#self.vline = ax.axvline(valinit, 0, 1, color='r', lw=1)
#self.valtext = ax.text(1.02, 0.5, valfmt % valinit,
self.valtext = self.ax.text(1.005, 0.5,self._value_format(valinit),
transform=self.ax.transAxes,
verticalalignment='center',
horizontalalignment='left')
def add_observer(self, obj):
"""
When the slider value is changed, call *func* with the new
slider position
A connection id is returned which can be used to disconnect
"""
self.observers[obj.name] = obj
def remove_observer(self, cid):
"""remove the observer with connection id *cid*"""
try:
del self.observers[cid]
except KeyError:
pass
def reset(self):
"""reset the slider to the initial value if needed"""
if (self.val != self.valinit):
self._set_val(self.valinit)
def _connect(self):
# 信号连接。
self.connect_event('button_press_event', self.on_event)
self.connect_event('button_release_event', self.on_event)
if self.drag_enabled:
self.connect_event('motion_notify_event', self.on_event)
def on_event(self, event):
"""update the slider position"""
if self.ignore(event):
return
if event.button != 1:
return
if event.name == 'button_press_event' and event.inaxes == self.ax:
self.drag_active = True
event.canvas.grab_mouse(self.ax)
if not self.drag_active:
return
elif event.name == 'button_press_event' and event.inaxes != self.ax:
self.drag_active = False
event.canvas.release_mouse(self.ax)
return
elif event.name == 'button_release_event' and event.inaxes == self.ax:
self.drag_active = False
event.canvas.release_mouse(self.ax)
self._update_observer(event)
return
self._update(event.xdata)
self._update_observer(event)
# 重绘
def _update(self, val, width=None):
if val <= self.valmin:
if not self.closedmin:
return
val = self.valmin
elif val >= self.valmax:
if not self.closedmax:
return
val = self.valmax
if self.slidermin is not None and val <= self.slidermin.val:
if not self.closedmin:
return
val = self.slidermin.val
if self.slidermax is not None and val >= self.slidermax.val:
if not self.closedmax:
return
val = self.slidermax.val
if width:
self.width = width
self._set_val(val)
def _set_val(self, val):
xy = self.poly.xy
xy[2] = val, 1
xy[3] = val, 0
self.val = val
self.poly.remove()
self.poly = self.ax.axvspan(val-self.width/2, val+self.width/2, 0, 1)
#self.poly.xy = xy
#self.valtext.set_text(self.valfmt % val)
self.valtext.set_text(self._value_format(val))
self.val = val
if not self.eventson:
return
def _update_observer(self, event):
""" 通知相关窗口更新数据 """
for name, obj in six.iteritems(self.observers):
try:
obj.on_slider(self.val, event)
except Exception as e:
six.print_(e)
class FrameWidget(AxesWidget):
"""
蜡烛线控件。
"""
def __init__(self, ax, name, wdlength, min_wdlength):
"""
Create a slider from *valmin* to *valmax* in axes *ax*
"""
AxesWidget.__init__(self, ax)
self.name = name
self.wdlength = wdlength
self.min_wdlength = min_wdlength
self.voffset = 0
self.connect()
self.plotters = { }
# 当前显示的范围。
#self.xmax = len(data)
#self.xmin = max(0, self.xmax-self.wdlength)
#self.ymax = np.max(data.high[self.xmin : self.xmax].values) + self.voffset
#self.ymin = np.min(data.low[self.xmin : self.xmax].values) - self.voffset
self.cnt = 0
self.observers = {}
def add_plotter(self, plotter, twinx):
""" 添加并绘制, 不允许重名的plotter """
if plotter.name in self.plotters:
raise
if not self.plotters:
twinx = False
if twinx:
twaxes = self.ax.twinx()
plotter.plot(twaxes)
plotter.ax = twaxes
plotter.twinx = True
else:
plotter.plot(self.ax)
plotter.ax = self.ax
plotter.twinx = False
self.plotters[plotter.name] = plotter
def plot_with_plotter(self, plotter_name, *args):
self.plotters[plotter_name].plot(self.ax, *args)
def set_ylim(self, w_left, w_right):
all_ymax = []
all_ymin = []
for plotter in six.itervalues(self.plotters):
if plotter.twinx:
continue
ymax, ymin = plotter.y_interval(w_left, w_right)
## @todo move ymax, ymin 计算到plot中去。
all_ymax.append(ymax)
all_ymin.append(ymin)
ymax = max(all_ymax)
ymin = min(all_ymin)
self._voffset = (ymax-ymin) / 10.0 # 画图显示的y轴留白。
ymax += self._voffset
ymin -= self._voffset
self.ax.set_ylim((ymin, ymax))
def on_slider(self, val, event):
#'''docstring for update(val)'''
## @TODO _set_ylim 分解到这里
pass
#val = int(val)
#self.xmax = val
#self.xmin = max(0, self.xmax-self.wdlength)
#self.ymax = np.max(self.data.high[val-self.wdlength : val].values) + self.voffset
#self.ymin = np.min(self.data.low[val-self.wdlength : val].values) - self.voffset
#self.ax.set_xlim((val-self.wdlength, val))
#self.ax.set_ylim((self.ymin, self.ymax))
def connect(self):
#self.ax.figure.canvas.mpl_connect('key_release_event', self.enter_axes)
pass
def _update(self, event):
"""update the slider position"""
self.update(event.xdata)
def add_observer(self, obj):
"""
When the slider value is changed, call *func* with the new
slider position
A connection id is returned which can be used to disconnect
"""
self.observers[obj.name] = obj
def disconnect(self, cid):
"""remove the observer with connection id *cid*"""
try:
del self.observers[cid]
except KeyError:
pass
def _update_observer(self, obname):
#"通知进度条改变宽度"
#for name, obj in six.iteritems(self.observers):
#if name == obname and obname == "slider":
#obj.update(obj.val, self.wdlength)
#break
pass
class MyLocator(mticker.MaxNLocator):
def __init__(self, *args, **kwargs):
mticker.MaxNLocator.__init__(self, *args, **kwargs)
def __call__(self, *args, **kwargs):
return mticker.MaxNLocator.__call__(self, *args, **kwargs)
class TechnicalWidget(object):
""" 多窗口控件 """
def __init__(self, fig, data, left=0.1, bottom=0.05, width=0.85, height=0.9,
parent=None):
""" 多窗口联动控件。
Args:
fig (Figure): matplotlib绘图容器。
data (DataFrame): [open, close, high, low]数据表。
"""
self.name = "MultiWidgets"
self._fig = fig
self._subwidgets = { }
self._cursor = None
self._cursor_axes_index = { }
self._hoffset = 1
self._left, self._width = left, width
self._bottom, self._height = bottom, height
self._slider_height = 0.1
self._bigger_picture_height = 0.3 # 鸟瞰图高度
self._all_axes = []
self.load_data(data)
self._cursor_axes = { }
def init_layout(self, w_width, *args):
# 布局参数
self._w_width_min = 50
self._w_width = w_width
self._init_widgets(*args)
self._connect()
self._cursor = MultiCursor(self._fig.canvas, self.axes,
color='r', lw=2, horizOn=False,
vertOn=True)
return self.axes
def load_data(self, data):
self._data = data
self._data_length = len(self._data)
@property
def axes(self):
return self._axes
def plot_text(self, name, ith_ax, x, y, text, color='black', size=10, rotation=0):
self.axes[ith_ax].text(x, y, text, color=color, fontsize=size, rotation=rotation)
def draw_widgets(self):
""" 显示控件 """
self._w_left = self._data_length - self._w_width
self._w_right = self._data_length
self._reset_auxiliary_widgets()
self._update_widgets()
def _reset_auxiliary_widgets(self):
if self._slider is None:
self._slider = Slider(self._slider_ax, "slider", '', 0, self._data_length-1,
self._data_length-1, self._data_length/50, "%d",
self._data.index)
self._slider.add_observer(self)
else:
self._slider.reinit( 0, self._data_length-1, self._data_length-1,
self._data_length/50, "%d", self._data.index)
if self._bigger_picture_plot:
self._bigger_picture_plot.pop(0).remove()
self._bigger_picture_plot = self._bigger_picture.plot(self._data['close'].values,
'b')
self._bigger_picture.set_ylim((min(self._data['low']), max(self._data['high'])))
self._bigger_picture.set_xlim((0, len(self._data['close'])))
self._slider_ax.xaxis.set_major_formatter(TimeFormatter(self._data.index,
fmt='%Y-%m-%d'))
def add_widget(self, ith_subwidget, widget, ymain=False, connect_slider=False):
""" 添加一个能接收消息事件的控件。
Args:
ith_subwidget (int.): 子窗口序号。
widget (AxesWidget): 控件。
Returns:
AxesWidget. widget
"""
# 对新创建的Axes做相应的处理
# 并且调整Cursor
for plotter in six.itervalues(widget.plotters):
if plotter.twinx:
plotter.ax.format_coord = self._format_coord
self.axes.append(plotter.ax)
#self._cursor_axes[ith_subwidget] = plotter.ax
self._cursor = MultiCursor(self._fig.canvas,
list(self._cursor_axes.values()),
color='r', lw=2, horizOn=False,
vertOn=True)
self._subwidgets[ith_subwidget] = widget
if connect_slider:
self._slider.add_observer(widget)
return widget
def on_slider(self, val, event):
""" 滑块事件处理。 """
if event.name == "button_press_event":
self._bigger_picture.set_zorder(1000)
self._slider_cursor = MultiCursor(self._fig.canvas,
[self._slider_ax, self._bigger_picture], color='y',
lw=2, horizOn=False, vertOn=True)
log.debug("on_press_event")
elif event.name == "button_release_event":
self._bigger_picture.set_zorder(0)
del self._slider_cursor
log.debug("on_release_event")
elif event.name == "motion_notify_event":
pass
# 遍历axes中的每个indicator,计算显示区间。
self._w_left = int(val)
self._w_right = self._w_left+self._w_width
if self._w_right >= self._data_length:
self._w_right = self._data_length - 1 + self._hoffset
self._w_left = self._w_right - self._w_width
self._update_widgets()
def on_press(self, event):
log.debug("button_press_event")
pass
def on_release(self, event):
pass
def on_motion(self, event):
#self._fig.canvas.draw()
pass
def _clear(self):
""""""
return
def on_keyrelease(self, event):
if event.key == u"down":
self._w_width += self._w_width/2
self._w_width = min(self._data_length, self._w_width)
elif event.key == u"up" :
self._w_width -= self._w_width/2
self._w_width= max(self._w_width, self._w_width_min)
elif event.key == u"super+up":
six.print_(event.key, "**", type(event.key) )
elif event.key == u"super+down":
six.print_(event.key, "**", type(event.key) )
# @TODO page upper down
middle = (self._w_left+self._w_right)/2
self._w_left = middle - self._w_width/2
self._w_right = middle + self._w_width/2
self._w_left = max(0, self._w_left)
self._w_right = min(self._data_length, self._w_right)
self._update_widgets()
def on_enter_axes(self, event):
#event.inaxes.patch.set_facecolor('yellow')
# 只有当前axes会闪烁。
if event.inaxes is self._slider_ax: #or event.inaxes is self._bigger_picture:
self._cursor = None
event.canvas.draw()
log.debug("on_enter_axes")
return
def on_leave_axes(self, event):
if event.inaxes is self._slider_ax:
# 进入后会创建_slider_cursor,离开后复原
axes = [self.axes[i] for i in six.itervalues(self._cursor_axes_index)]
#axes = list(reversed(axes)) # 很奇怪,如果没有按顺序给出,显示会有问题。
self._cursor = MultiCursor(self._fig.canvas, axes, color='r', lw=2, horizOn=False, vertOn=True)
event.canvas.draw()
log.debug("on_leave_axes")
def _connect(self):
"""
matplotlib信号连接。
"""
self.cidpress = self._fig.canvas.mpl_connect( "button_press_event", self.on_press)
self.cidrelease = self._fig.canvas.mpl_connect( "button_release_event", self.on_release)
self.cidmotion = self._fig.canvas.mpl_connect( "motion_notify_event", self.on_motion)
self._fig.canvas.mpl_connect('axes_enter_event', self.on_enter_axes)
self._fig.canvas.mpl_connect('axes_leave_event', self.on_leave_axes)
self._fig.canvas.mpl_connect('key_release_event', self.on_keyrelease)
#def _disconnect(self):
#self._fig.canvas.mpl_disconnect(self.cidmotion)
#self._fig.canvas.mpl_disconnect(self.cidrelease)
#self._fig.canvas.mpl_disconnect(self.cidpress)
def _init_widgets(self, *args):
self._slidder_lower = self._bottom
self._slidder_upper = self._bottom + self._slider_height
self._bigger_picture_lower = self._slidder_upper
self._slider_ax = self._fig.add_axes([self._left, self._slidder_lower, self._width,
self._slider_height])
self._bigger_picture = self._fig.add_axes([self._left, self._bigger_picture_lower,
self._width, self._bigger_picture_height],
zorder = 0, frameon=False,
#sharex=self._slider_ax,
alpha = '0.1' )
self._bigger_picture.set_xticklabels([]);
self._bigger_picture.set_xticks([])
self._bigger_picture.set_yticks([])
self._all_axes = [self._slider_ax, self._bigger_picture]
self._slider = None
self._bigger_picture_plot = None
args = list(reversed(args))
# 默认子窗口数量为1
if len(args) == 0:
args = (1,)
total_units = sum(args)
unit = (self._bottom + self._height - self._slidder_upper) / total_units
bottom = self._slidder_upper
user_axes = []
first_user_axes = None
for i, ratio in enumerate(args):
rect = [self._left, bottom, self._width, unit * ratio]
if i > 0:
# 共享x轴
ax = self._fig.add_axes(rect, sharex=first_user_axes) #facecolor=axescolor)
self._all_axes.append(ax)
else:
first_user_axes = self._fig.add_axes(rect)
self._all_axes.append(first_user_axes)
user_axes = self._all_axes[2:]
bottom += unit * ratio
self._axes = list(reversed(user_axes))
map(lambda x: x.grid(True), self._axes)
map(lambda x: x.set_xticklabels([]), self._axes[1:])
for ax in self.axes:
ax.get_yaxis().get_major_formatter().set_useOffset(False)
# ax.get_yaxis().get_major_formatter().set_scientific(False)
for i, ax in enumerate(self.axes):
ax.format_coord = self._format_coord
self._cursor_axes[i] = ax
delta = (self._data.index[1] - self._data.index[0])
self.axes[0].xaxis.set_major_formatter(TimeFormatter(self._data.index, delta))
self.axes[0].set_xticks(self._xticks_to_display(0, self._data_length, delta));
for ax in self.axes[0:-1]:
[label.set_visible(False) for label in ax.get_xticklabels()]
for i in range(0, len(self.axes)):
self._cursor_axes_index[i] = i
def _update_widgets(self):
""" 改变可视区域, 在坐标移动后被调用。"""
self.axes[0].set_xlim((int(self._w_left), int(self._w_right)))
self._set_ylim(int(self._w_left), int(self._w_right))
self._fig.canvas.draw()
def _set_ylim(self, w_left, w_right):
""" 设置当前显示窗口的y轴范围。
"""
for subwidget in six.itervalues(self._subwidgets):
subwidget.set_ylim(w_left, w_right)
def _format_coord(self, x, y):
""" 状态栏信息显示 """
index = x
f = x % 1
index = x-f if f < 0.5 else min(x-f+1, len(self._data['open']) - 1)
delta = (self._data.index[1] - self._data.index[0])
fmt = slider_strtime_format(delta)
index = int(index)
## @note 字符串太长会引起闪烁
return "[dt=%s o=%.2f c=%.2f h=%.2f l=%.2f]" % (
self._data.index[index].strftime(fmt),
self._data['open'][index],
self._data['close'][index],
self._data['high'][index],
self._data['low'][index])
def _xticks_to_display(self, start, end, delta):
xticks = []
for i in range(start, end):
if i >= 1:
if delta.days >= 1:
if self._data.index[i].month != self._data.index[i-1].month:
xticks.append(i)
elif delta.seconds == 60:
# 一分钟的以小时为显示单位
if self._data.index[i].hour != self._data.index[i-1].hour and \
self._data.index[i].day == self._data.index[i-1].day:
xticks.append(i)
else:
if self._data.index[i].day != self._data.index[i-1].day:
# 其它日内以天为显示单位
xticks.append(i)
else:
xticks.append(0)
return xticks
class TimeFormatter(Formatter):
# 分类 --format
def __init__(self, dates, delta=None, fmt='%Y-%m-%d %H:%M'):
self.dates = dates
self.fmt = self._strtime_format(delta) if delta else fmt
def __call__(self, x, pos=0):
'Return the label for time x at position pos'
ind = int(round(x))
if ind>=len(self.dates) or ind<0: return ''
return self.dates[ind].strftime(self.fmt)
def _strtime_format(self, delta):
if delta.days >= 1:
return '%Y-%m'
elif delta.seconds == 60:
return '%m-%d %H:%M'
else:
# 日内其它分钟
return '%m-%d'