expr-tracker 0.1.7__tar.gz → 0.1.8__tar.gz

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.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: expr_tracker
3
- Version: 0.1.7
3
+ Version: 0.1.8
4
4
  Summary: Add your description here
5
5
  Author-email: HSPK <whxway@whu.edu.cn>
6
6
  Requires-Python: >=3.10
@@ -321,6 +321,30 @@ def jsonable_encoder(
321
321
  if isinstance(obj, classes_tuple):
322
322
  return encoder(obj)
323
323
 
324
+ # numpy 标量/数组、torch tensor 等既不是 int/float 的子类,也无法被 dict()/vars()
325
+ # 处理,但都提供 tolist()/item(),转换后再递归编码(如 datetime64 -> datetime)。
326
+ for method_name in ("tolist", "item"):
327
+ method = getattr(obj, method_name, None)
328
+ if not callable(method):
329
+ continue
330
+ try:
331
+ converted = method()
332
+ except Exception:
333
+ continue
334
+ if converted is obj:
335
+ break
336
+ return jsonable_encoder(
337
+ converted,
338
+ include=include,
339
+ exclude=exclude,
340
+ by_alias=by_alias,
341
+ exclude_unset=exclude_unset,
342
+ exclude_defaults=exclude_defaults,
343
+ exclude_none=exclude_none,
344
+ custom_encoder=custom_encoder,
345
+ sqlalchemy_safe=sqlalchemy_safe,
346
+ )
347
+
324
348
  try:
325
349
  data = dict(obj)
326
350
  except Exception as e:
@@ -0,0 +1,350 @@
1
+ import atexit
2
+ import json
3
+ import os
4
+ import threading
5
+ import time
6
+ from pathlib import Path
7
+
8
+ from loguru import logger
9
+
10
+ from expr_tracker.encoders import jsonable_encoder
11
+
12
+ DEFAULT_BUFFER_SIZE = 50
13
+ DEFAULT_BUFFER_INTERVAL = 1.0
14
+ DEFAULT_MAX_BUFFER_SECONDS = 5.0
15
+ # 写盘持续失败时 buffer 的最大长度,超出后丢弃最旧的记录,避免内存无限增长
16
+ DEFAULT_MAX_PENDING_RECORDS = 100_000
17
+ MAX_FALLBACK_REPR_LENGTH = 512
18
+
19
+ # 与 jsonlines 默认(非 compact)writer 的输出保持一致
20
+ _LINE_ENCODER = json.JSONEncoder(ensure_ascii=False)
21
+
22
+
23
+ def _fallback_repr(value) -> str:
24
+ """无法 JSON 序列化时的兜底表示,保证永远不抛异常"""
25
+ try:
26
+ text = repr(value)
27
+ except Exception:
28
+ return f"<unrepresentable {type(value).__name__}>"
29
+ if len(text) > MAX_FALLBACK_REPR_LENGTH:
30
+ text = text[:MAX_FALLBACK_REPR_LENGTH] + "..."
31
+ return text
32
+
33
+
34
+ def _encode_key(key) -> str:
35
+ if isinstance(key, str):
36
+ return key
37
+ try:
38
+ encoded = jsonable_encoder(key)
39
+ except Exception:
40
+ encoded = None
41
+ if isinstance(encoded, str):
42
+ return encoded
43
+ if isinstance(encoded, (int, float, bool)) or encoded is None:
44
+ return str(encoded)
45
+ return _fallback_repr(key)
46
+
47
+
48
+ class JsonlTracker:
49
+ def __init__(self):
50
+ self.buffer = []
51
+ self.buffer_size = DEFAULT_BUFFER_SIZE
52
+ self.buffer_interval = DEFAULT_BUFFER_INTERVAL
53
+ self.max_buffer_seconds = DEFAULT_MAX_BUFFER_SECONDS
54
+ self.max_pending_records = DEFAULT_MAX_PENDING_RECORDS
55
+ self.log_fp = None
56
+ self._lock = threading.RLock()
57
+ # 只用于串行化磁盘写入,不阻塞 log() 写 buffer
58
+ self._write_lock = threading.Lock()
59
+ self._last_log_time = None
60
+ self._first_buffered_time = None
61
+ self._flush_timer = None
62
+ self._atexit_registered = False
63
+ self._warned_metric_keys = set()
64
+
65
+ def init(
66
+ self,
67
+ project: str,
68
+ name: str | None = None,
69
+ config: dict | None = None,
70
+ dir: str | None = None,
71
+ print_to_screen: bool = False,
72
+ print_handle=print,
73
+ buffer_size: int = DEFAULT_BUFFER_SIZE,
74
+ buffer_interval: float | None = DEFAULT_BUFFER_INTERVAL,
75
+ max_buffer_seconds: float | None = DEFAULT_MAX_BUFFER_SECONDS,
76
+ max_pending_records: int = DEFAULT_MAX_PENDING_RECORDS,
77
+ **kwargs,
78
+ ):
79
+ """初始化 jsonl backend。
80
+
81
+ 缓冲策略(按 log 频率自适应):
82
+ - ``buffer_size``: buffer 中记录数达到该值立即写盘。
83
+ - ``buffer_interval``: 相邻两次 ``log()`` 的间隔 >= 该值时认为不是高频写入,
84
+ 直接写盘(低延迟);小于该值才认为是高频写入,先攒在内存里。
85
+ 设为 ``None`` 表示关闭该判断(只按 buffer_size 攒批)。
86
+ - ``max_buffer_seconds``: 记录在 buffer 中的最长停留时间,超时后由后台定时器
87
+ 强制写盘,避免高频写入突然停止时数据长期滞留内存。设为 ``None`` 关闭。
88
+ - ``max_pending_records``: 写盘持续失败(磁盘满、挂载断开等)时 buffer 的上限,
89
+ 超出后丢弃最旧的记录,避免 OOM。
90
+ """
91
+ self.project = project
92
+ if name is None:
93
+ name = time.strftime("run-%Y%m%d-%H%M%S")
94
+ logger.warning(f"No run name provided, using generated name: {name}")
95
+ self.name = name
96
+ if dir is None:
97
+ dir = "./tracker/jsonl"
98
+ self.log_dir = Path(dir) / self.project / self.name
99
+ self.config_fp = self.log_dir / "config.json"
100
+ self.log_fp = self.log_dir / "metrics.jsonl"
101
+
102
+ # 初始化 Buffer 配置
103
+ self._cancel_timer()
104
+ with self._lock:
105
+ self.buffer = []
106
+ self.buffer_size = max(1, int(buffer_size))
107
+ self.buffer_interval = (
108
+ None if buffer_interval is None else float(buffer_interval)
109
+ )
110
+ self.max_buffer_seconds = (
111
+ None if max_buffer_seconds is None else float(max_buffer_seconds)
112
+ )
113
+ self.max_pending_records = max(
114
+ self.buffer_size, int(max_pending_records or 0)
115
+ )
116
+ self._last_log_time = None
117
+ self._first_buffered_time = None
118
+ self._warned_metric_keys = set()
119
+
120
+ self.log_dir.mkdir(parents=True, exist_ok=True)
121
+
122
+ if self.config_fp.exists():
123
+ logger.warning(
124
+ f"Config file {self.config_fp} already exists. It will be overwritten."
125
+ )
126
+
127
+ if config is not None:
128
+ # Config 通常只写一次,直接写入即可
129
+ try:
130
+ with open(self.config_fp, "w", encoding="utf-8") as f:
131
+ json.dump(
132
+ self._encode_mapping(config, kind="config"),
133
+ f,
134
+ indent=4,
135
+ ensure_ascii=False,
136
+ )
137
+ except Exception as e:
138
+ logger.error(f"Failed to write config to {self.config_fp}: {e}")
139
+
140
+ self.print_to_screen = print_to_screen
141
+ self.print_handle = print_handle
142
+
143
+ # 优化:流式计算行数,避免一次性加载大文件到内存 (对 BlobFuse 友好)
144
+ self.current_step = 0
145
+ if self.log_fp.exists():
146
+ try:
147
+ with open(self.log_fp, "rb") as f:
148
+ self.current_step = sum(1 for _ in f)
149
+ except Exception as e:
150
+ logger.warning(f"Could not count existing lines in {self.log_fp}: {e}")
151
+
152
+ # 注册退出钩子:确保程序意外终止时也能写入剩余数据
153
+ if not self._atexit_registered:
154
+ atexit.register(self.flush)
155
+ self._atexit_registered = True
156
+
157
+ def _encode_mapping(self, metrics: dict | None, kind: str = "metric") -> dict:
158
+ """把 metrics 转成可 JSON 序列化的 dict,绝不因单个坏值而丢掉整条记录。
159
+
160
+ numpy/torch 标量等由 ``jsonable_encoder`` 统一处理;真正无法编码的值降级为
161
+ ``repr``,并按 key 只告警一次,避免高频训练循环刷屏。
162
+ """
163
+ if not metrics:
164
+ return {}
165
+ if not isinstance(metrics, dict):
166
+ return {"value": self._encode_value("value", metrics, kind)}
167
+ try:
168
+ encoded = jsonable_encoder(metrics)
169
+ if isinstance(encoded, dict):
170
+ return encoded
171
+ except Exception:
172
+ pass
173
+ # 整体编码失败:逐字段降级,保留其余可序列化的字段
174
+ encoded = {}
175
+ for key, value in metrics.items():
176
+ encoded_key = _encode_key(key)
177
+ encoded[encoded_key] = self._encode_value(encoded_key, value, kind)
178
+ return encoded
179
+
180
+ def _encode_value(self, key: str, value, kind: str = "metric"):
181
+ try:
182
+ return jsonable_encoder(value)
183
+ except Exception as e:
184
+ if key not in self._warned_metric_keys:
185
+ self._warned_metric_keys.add(key)
186
+ logger.warning(
187
+ f"{kind.capitalize()} {key!r} of type {type(value).__name__} is not "
188
+ f"JSON serializable ({e}); falling back to repr()."
189
+ )
190
+ return _fallback_repr(value)
191
+
192
+ def log(self, metrics: dict, step: int | None = None):
193
+ now = time.monotonic()
194
+ # 在锁外编码:大对象的转换不阻塞其他线程,同时保证坏值不会进入 buffer
195
+ encoded_metrics = self._encode_mapping(metrics)
196
+
197
+ with self._lock:
198
+ if step is not None:
199
+ self.current_step = step
200
+
201
+ record = {"_step": self.current_step, **encoded_metrics}
202
+
203
+ # 1. 写入内存 Buffer
204
+ self.buffer.append(record)
205
+ if self._first_buffered_time is None:
206
+ self._first_buffered_time = now
207
+
208
+ # 2. 根据「距离上一次 log 的时间间隔」判断是否为高频写入
209
+ interval = None if self._last_log_time is None else now - self._last_log_time
210
+ self._last_log_time = now
211
+
212
+ should_flush = self._should_flush(now, interval)
213
+ self.current_step += 1
214
+
215
+ # 3. 屏幕打印(放在锁外,避免 print handle 阻塞其他线程)
216
+ if self.print_to_screen:
217
+ try:
218
+ self.print_handle(f"{record}")
219
+ except Exception as e:
220
+ logger.warning(f"Failed to print metrics to screen: {e}")
221
+
222
+ # 4. 立即写盘,或安排一次超时写盘
223
+ if should_flush:
224
+ self.flush()
225
+ else:
226
+ self._schedule_timer()
227
+
228
+ def _should_flush(self, now: float, interval: float | None) -> bool:
229
+ """判断当前是否应该立即写盘(需在持锁状态下调用)"""
230
+ # Buffer 已满
231
+ if len(self.buffer) >= self.buffer_size:
232
+ return True
233
+ # 首次 log:直接落盘,尽快产生文件内容
234
+ if interval is None:
235
+ return True
236
+ # 低频写入:距离上次 log 已经过去足够久,没必要继续攒批
237
+ if self.buffer_interval is not None and interval >= self.buffer_interval:
238
+ return True
239
+ # 高频写入,但最早的记录已在内存中停留过久
240
+ return (
241
+ self.max_buffer_seconds is not None
242
+ and self._first_buffered_time is not None
243
+ and now - self._first_buffered_time >= self.max_buffer_seconds
244
+ )
245
+
246
+ def _schedule_timer(self):
247
+ """为 buffer 中最早的记录安排一次超时写盘"""
248
+ if self.max_buffer_seconds is None:
249
+ return
250
+ with self._lock:
251
+ if self._flush_timer is not None or not self.buffer:
252
+ return
253
+ now = time.monotonic()
254
+ elapsed = now - (self._first_buffered_time or now)
255
+ delay = max(0.0, self.max_buffer_seconds - elapsed)
256
+ timer = threading.Timer(delay, self._on_timer)
257
+ timer.daemon = True
258
+ self._flush_timer = timer
259
+ timer.start()
260
+
261
+ def _on_timer(self):
262
+ with self._lock:
263
+ self._flush_timer = None
264
+ self.flush()
265
+
266
+ def _cancel_timer(self):
267
+ with self._lock:
268
+ timer = self._flush_timer
269
+ self._flush_timer = None
270
+ if timer is not None:
271
+ timer.cancel()
272
+
273
+ def flush(self):
274
+ """强制将内存中的 Buffer 写入磁盘"""
275
+ self._cancel_timer()
276
+
277
+ with self._lock:
278
+ self._first_buffered_time = None
279
+ if not self.buffer:
280
+ return
281
+ log_fp = self.log_fp
282
+ if log_fp is None:
283
+ logger.error(
284
+ f"JsonlTracker is not initialized, dropping {len(self.buffer)} "
285
+ "buffered records. Call init() before log()."
286
+ )
287
+ self.buffer = []
288
+ return
289
+ records, self.buffer = self.buffer, []
290
+
291
+ # 先在内存里序列化:无法序列化的记录直接丢弃,否则它会永远卡在 buffer 里,
292
+ # 导致后续所有指标都写不进去
293
+ lines, pending = [], []
294
+ for record in records:
295
+ try:
296
+ lines.append(_LINE_ENCODER.encode(record) + "\n")
297
+ pending.append(record)
298
+ except Exception as e:
299
+ logger.error(f"Dropping record that cannot be serialized: {e}")
300
+ if not lines:
301
+ return
302
+
303
+ size_before = None
304
+ with self._write_lock:
305
+ try:
306
+ # 确保目录存在 (防止运行时目录被删)
307
+ if not log_fp.parent.exists():
308
+ log_fp.parent.mkdir(parents=True, exist_ok=True)
309
+ size_before = log_fp.stat().st_size if log_fp.exists() else 0
310
+ # 一次性写入整批,避免中途失败留下「已写一半」的批次
311
+ with open(log_fp, "a", encoding="utf-8") as f:
312
+ f.write("".join(lines))
313
+ except Exception as e:
314
+ logger.error(f"Failed to flush metrics to {log_fp}: {e}")
315
+ self._truncate_partial_write(log_fp, size_before)
316
+ # 写入失败:放回 buffer 头部,等待下次 flush 重试
317
+ self._requeue(pending)
318
+ self._schedule_timer()
319
+
320
+ @staticmethod
321
+ def _truncate_partial_write(log_fp: Path, size_before: int | None):
322
+ """回滚半截写入,保证重试时不会产生重复或损坏的行"""
323
+ if size_before is None:
324
+ return
325
+ try:
326
+ if log_fp.exists() and log_fp.stat().st_size > size_before:
327
+ os.truncate(log_fp, size_before)
328
+ except Exception as e:
329
+ logger.warning(f"Could not roll back partial write on {log_fp}: {e}")
330
+
331
+ def _requeue(self, records: list):
332
+ with self._lock:
333
+ self.buffer[:0] = records
334
+ overflow = len(self.buffer) - self.max_pending_records
335
+ if overflow > 0:
336
+ del self.buffer[:overflow]
337
+ logger.error(
338
+ f"Metrics buffer exceeded {self.max_pending_records} records, "
339
+ f"dropped {overflow} oldest records."
340
+ )
341
+ self._first_buffered_time = time.monotonic()
342
+
343
+ def finish(self):
344
+ """结束时显式调用"""
345
+ self._cancel_timer()
346
+ self.flush()
347
+ # 如果手动调用了 finish,取消 atexit 注册,防止重复调用
348
+ if self._atexit_registered:
349
+ atexit.unregister(self.flush)
350
+ self._atexit_registered = False
@@ -105,6 +105,110 @@ def test_concurrent_logging_does_not_lose_records():
105
105
  assert _count_lines(tracker.log_fp) == 800
106
106
 
107
107
 
108
+ def test_numpy_values_are_serialized():
109
+ """numpy 标量/数组不是 int/float 子类,必须被编码后才能写入 jsonl"""
110
+ try:
111
+ import numpy as np
112
+ except ImportError: # numpy 是可选依赖
113
+ return
114
+ tracker = _new_tracker("numpy", buffer_size=1)
115
+ tracker.log(
116
+ {
117
+ "int32": np.int32(7),
118
+ "float32": np.float32(0.5),
119
+ "bool": np.bool_(True),
120
+ "array": np.array([[1, 2], [3, 4]]),
121
+ "date": np.datetime64("2024-01-01"),
122
+ }
123
+ )
124
+ tracker.finish()
125
+
126
+ with jsonlines.open(tracker.log_fp) as reader:
127
+ record = next(iter(reader))
128
+ assert record["int32"] == 7
129
+ assert record["float32"] == 0.5
130
+ assert record["bool"] is True
131
+ assert record["array"] == [[1, 2], [3, 4]]
132
+ assert record["date"] == "2024-01-01"
133
+ assert tracker.buffer == []
134
+
135
+
136
+ def test_unserializable_value_does_not_block_other_records():
137
+ """无法序列化的值降级为 repr,不能污染 buffer 或丢掉后续记录"""
138
+
139
+ class Weird:
140
+ __slots__ = ()
141
+
142
+ def __repr__(self):
143
+ return "<weird>"
144
+
145
+ tracker = _new_tracker("poison", buffer_size=2, buffer_interval=None)
146
+ tracker.log({"v": 0})
147
+ tracker.log({"bad": Weird(), "good": 1})
148
+ tracker.log({"v": 2})
149
+ tracker.finish()
150
+
151
+ with jsonlines.open(tracker.log_fp) as reader:
152
+ records = list(reader)
153
+ assert [record["_step"] for record in records] == [0, 1, 2]
154
+ assert records[1] == {"_step": 1, "bad": "<weird>", "good": 1}
155
+ assert tracker.buffer == []
156
+
157
+
158
+ def test_write_failure_retries_without_duplicates():
159
+ """写盘失败时整批回到 buffer,恢复后补写且不产生重复行"""
160
+ tracker = _new_tracker(
161
+ "io_fail", buffer_size=2, buffer_interval=None, max_buffer_seconds=None
162
+ )
163
+ tracker.log({"v": 0})
164
+ assert _count_lines(tracker.log_fp) == 1
165
+
166
+ good_fp = tracker.log_fp
167
+ blocked = tracker.log_dir / "blocked"
168
+ blocked.mkdir()
169
+ tracker.log_fp = blocked # 写入目录必然失败
170
+
171
+ tracker.log({"v": 1})
172
+ tracker.log({"v": 2})
173
+ assert len(tracker.buffer) == 2
174
+ assert _count_lines(good_fp) == 1
175
+
176
+ tracker.log_fp = good_fp
177
+ tracker.finish()
178
+
179
+ with jsonlines.open(good_fp) as reader:
180
+ assert [record["_step"] for record in reader] == [0, 1, 2]
181
+
182
+
183
+ def test_buffer_is_capped_when_writes_keep_failing():
184
+ """持续写失败时 buffer 不应无限增长"""
185
+ tracker = _new_tracker(
186
+ "capped",
187
+ buffer_size=2,
188
+ buffer_interval=None,
189
+ max_buffer_seconds=None,
190
+ max_pending_records=4,
191
+ )
192
+ blocked = tracker.log_dir / "blocked"
193
+ blocked.mkdir()
194
+ tracker.log_fp = blocked
195
+
196
+ for i in range(20):
197
+ tracker.log({"v": i})
198
+ assert len(tracker.buffer) <= 4
199
+ # 保留的是最新的记录
200
+ assert tracker.buffer[-1]["v"] == 19
201
+
202
+
203
+ def test_init_without_name_generates_one():
204
+ tracker = JsonlTracker()
205
+ tracker.init(project="p", dir=tempfile.mkdtemp())
206
+ assert tracker.name
207
+ tracker.log({"v": 0})
208
+ tracker.finish()
209
+ assert _count_lines(tracker.log_fp) == 1
210
+
211
+
108
212
  if __name__ == "__main__":
109
213
  for _name, _fn in sorted(globals().items()):
110
214
  if _name.startswith("test_") and callable(_fn):
@@ -1,209 +0,0 @@
1
- import atexit
2
- import json
3
- import threading
4
- import time
5
- from pathlib import Path
6
-
7
- import jsonlines
8
- from loguru import logger
9
-
10
- from expr_tracker.encoders import jsonable_encoder
11
-
12
- DEFAULT_BUFFER_SIZE = 50
13
- DEFAULT_BUFFER_INTERVAL = 1.0
14
- DEFAULT_MAX_BUFFER_SECONDS = 5.0
15
-
16
-
17
- class JsonlTracker:
18
- def __init__(self):
19
- self.buffer = []
20
- self.buffer_size = DEFAULT_BUFFER_SIZE
21
- self.buffer_interval = DEFAULT_BUFFER_INTERVAL
22
- self.max_buffer_seconds = DEFAULT_MAX_BUFFER_SECONDS
23
- self.log_fp = None
24
- self._lock = threading.RLock()
25
- self._last_log_time = None
26
- self._first_buffered_time = None
27
- self._flush_timer = None
28
-
29
- def init(
30
- self,
31
- project: str,
32
- name: str | None = None,
33
- config: dict | None = None,
34
- dir: str | None = None,
35
- print_to_screen: bool = False,
36
- print_handle=print,
37
- buffer_size: int = DEFAULT_BUFFER_SIZE,
38
- buffer_interval: float | None = DEFAULT_BUFFER_INTERVAL,
39
- max_buffer_seconds: float | None = DEFAULT_MAX_BUFFER_SECONDS,
40
- **kwargs,
41
- ):
42
- """初始化 jsonl backend。
43
-
44
- 缓冲策略(按 log 频率自适应):
45
- - ``buffer_size``: buffer 中记录数达到该值立即写盘。
46
- - ``buffer_interval``: 相邻两次 ``log()`` 的间隔 >= 该值时认为不是高频写入,
47
- 直接写盘(低延迟);小于该值才认为是高频写入,先攒在内存里。
48
- 设为 ``None`` 表示关闭该判断(只按 buffer_size 攒批)。
49
- - ``max_buffer_seconds``: 记录在 buffer 中的最长停留时间,超时后由后台定时器
50
- 强制写盘,避免高频写入突然停止时数据长期滞留内存。设为 ``None`` 关闭。
51
- """
52
- self.project = project
53
- self.name = name
54
- if dir is None:
55
- dir = "./tracker/jsonl"
56
- self.log_dir = Path(dir) / self.project / self.name
57
- self.config_fp = self.log_dir / "config.json"
58
- self.log_fp = self.log_dir / "metrics.jsonl"
59
-
60
- # 初始化 Buffer 配置
61
- self._cancel_timer()
62
- with self._lock:
63
- self.buffer = []
64
- self.buffer_size = max(1, int(buffer_size))
65
- self.buffer_interval = (
66
- None if buffer_interval is None else float(buffer_interval)
67
- )
68
- self.max_buffer_seconds = (
69
- None if max_buffer_seconds is None else float(max_buffer_seconds)
70
- )
71
- self._last_log_time = None
72
- self._first_buffered_time = None
73
-
74
- self.log_dir.mkdir(parents=True, exist_ok=True)
75
-
76
- if self.config_fp.exists():
77
- logger.warning(
78
- f"Config file {self.config_fp} already exists. It will be overwritten."
79
- )
80
-
81
- if config is not None:
82
- # Config 通常只写一次,直接写入即可
83
- with open(self.config_fp, "w") as f:
84
- json.dump(jsonable_encoder(config), f, indent=4)
85
-
86
- self.print_to_screen = print_to_screen
87
- self.print_handle = print_handle
88
-
89
- # 优化:流式计算行数,避免一次性加载大文件到内存 (对 BlobFuse 友好)
90
- self.current_step = 0
91
- if self.log_fp.exists():
92
- try:
93
- with open(self.log_fp, "rb") as f:
94
- self.current_step = sum(1 for _ in f)
95
- except Exception as e:
96
- logger.warning(f"Could not count existing lines in {self.log_fp}: {e}")
97
-
98
- # 注册退出钩子:确保程序意外终止时也能写入剩余数据
99
- atexit.register(self.flush)
100
-
101
- def log(self, metrics: dict, step: int | None = None):
102
- now = time.monotonic()
103
-
104
- with self._lock:
105
- if step is not None:
106
- self.current_step = step
107
-
108
- record = {"_step": self.current_step, **metrics}
109
-
110
- # 1. 写入内存 Buffer
111
- self.buffer.append(record)
112
- if self._first_buffered_time is None:
113
- self._first_buffered_time = now
114
-
115
- # 2. 根据「距离上一次 log 的时间间隔」判断是否为高频写入
116
- interval = None if self._last_log_time is None else now - self._last_log_time
117
- self._last_log_time = now
118
-
119
- should_flush = self._should_flush(now, interval)
120
- self.current_step += 1
121
-
122
- # 3. 屏幕打印(放在锁外,避免 print handle 阻塞其他线程)
123
- if self.print_to_screen:
124
- self.print_handle(f"{record}")
125
-
126
- # 4. 立即写盘,或安排一次超时写盘
127
- if should_flush:
128
- self.flush()
129
- else:
130
- self._schedule_timer()
131
-
132
- def _should_flush(self, now: float, interval: float | None) -> bool:
133
- """判断当前是否应该立即写盘(需在持锁状态下调用)"""
134
- # Buffer 已满
135
- if len(self.buffer) >= self.buffer_size:
136
- return True
137
- # 首次 log:直接落盘,尽快产生文件内容
138
- if interval is None:
139
- return True
140
- # 低频写入:距离上次 log 已经过去足够久,没必要继续攒批
141
- if self.buffer_interval is not None and interval >= self.buffer_interval:
142
- return True
143
- # 高频写入,但最早的记录已在内存中停留过久
144
- return (
145
- self.max_buffer_seconds is not None
146
- and self._first_buffered_time is not None
147
- and now - self._first_buffered_time >= self.max_buffer_seconds
148
- )
149
-
150
- def _schedule_timer(self):
151
- """为 buffer 中最早的记录安排一次超时写盘"""
152
- if self.max_buffer_seconds is None:
153
- return
154
- with self._lock:
155
- if self._flush_timer is not None or not self.buffer:
156
- return
157
- now = time.monotonic()
158
- elapsed = now - (self._first_buffered_time or now)
159
- delay = max(0.0, self.max_buffer_seconds - elapsed)
160
- timer = threading.Timer(delay, self._on_timer)
161
- timer.daemon = True
162
- self._flush_timer = timer
163
- timer.start()
164
-
165
- def _on_timer(self):
166
- with self._lock:
167
- self._flush_timer = None
168
- self.flush()
169
-
170
- def _cancel_timer(self):
171
- with self._lock:
172
- timer = self._flush_timer
173
- self._flush_timer = None
174
- if timer is not None:
175
- timer.cancel()
176
-
177
- def flush(self):
178
- """强制将内存中的 Buffer 写入磁盘"""
179
- self._cancel_timer()
180
-
181
- with self._lock:
182
- self._first_buffered_time = None
183
- if not self.buffer:
184
- return
185
- records, self.buffer = self.buffer, []
186
- log_fp = self.log_fp
187
-
188
- # 确保目录存在 (防止运行时目录被删)
189
- if log_fp and not log_fp.parent.exists():
190
- log_fp.parent.mkdir(parents=True, exist_ok=True)
191
-
192
- try:
193
- # 批量追加写入
194
- with jsonlines.open(log_fp, mode="a") as writer:
195
- writer.write_all(records)
196
- except Exception as e:
197
- logger.error(f"Failed to flush metrics to {log_fp}: {e}")
198
- # 写入失败:放回 buffer 头部,等待下次 flush 重试
199
- with self._lock:
200
- self.buffer[:0] = records
201
- self._first_buffered_time = time.monotonic()
202
- self._schedule_timer()
203
-
204
- def finish(self):
205
- """结束时显式调用"""
206
- self._cancel_timer()
207
- self.flush()
208
- # 如果手动调用了 finish,取消 atexit 注册,防止重复调用
209
- atexit.unregister(self.flush)
File without changes
File without changes
File without changes