File size: 4,412 Bytes
faea13d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
84fbb3b
 
 
 
 
faea13d
84fbb3b
faea13d
 
 
 
 
 
84fbb3b
 
faea13d
 
 
 
 
84fbb3b
 
 
 
 
 
 
 
 
 
 
 
 
 
faea13d
 
 
 
 
 
 
 
 
 
 
 
84fbb3b
 
faea13d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""實測這台機器的 SSD 讀取能力。

**換機器一定要重測這個數字**,tok/s 幾乎完全由它決定(dense/每 token 重讀整個權重)。
輸出 JSON 給 validate/io-probe.json。

    python3 tools/io_probe.py --out validate/io-probe.json --size-mib 1024
"""
import argparse
import json
import os
import sys
import time

MIB = 1024 * 1024


def drop_cache(fd: int, size: int) -> None:
    os.posix_fadvise(fd, 0, size, 4)  # POSIX_FADV_DONTNEED
    os.fsync(fd)


def read_seq(path: str, size: int, threads: int) -> dict:
    """threads 條執行緒各自讀不同區段,量「有效頻寬」。

    迴圈裡不能把 bytes 寫進 bytes 物件(Windows 沒事,但這裡要可攜);
    一律用 os.pread 拿回讀到的長度。
    """
    import threading

    per = size // threads
    results = [0] * threads
    fds = [os.open(path, os.O_RDONLY) for _ in range(threads)]
    for fd in fds:
        drop_cache(fd, size)

    errors: list[str] = []

    def work(i: int):
        fd = fds[i]
        off = i * per
        left = per
        got = 0
        try:
            while left > 0:
                want = min(4 * MIB, left)
                # os.pread 回傳的是**資料本身**,不是讀到的位元組數。
                # 寫成 `got += n` 會得到 "int + bytes" 的 TypeError,
                # 而執行緒裡的例外不會讓主程式失敗 —— 結果就是
                # 「bytes: 0、0.015 秒、mb_per_s 0.0」這種完全沒有 I/O 的假象。
                buf = os.pread(fd, want, off + got)
                if not buf:
                    break
                got += len(buf)
                left -= len(buf)
        except OSError as e:
            errors.append(str(e))
        results[i] = got

    ts = [threading.Thread(target=work, args=(i,)) for i in range(threads)]
    t0 = time.time()
    for t in ts:
        t.start()
    for t in ts:
        t.join()
    dt = time.time() - t0
    for fd in fds:
        os.close(fd)
    total = sum(results)
    if errors:
        print(f"[io_probe] thread errors: {errors[:3]}", file=sys.stderr)
    return {"threads": threads, "bytes": total, "seconds": round(dt, 3),
            "mb_per_s": round(total / MIB / dt, 1) if dt > 0 else None}


def read_rand4k(path: str, count: int) -> dict:
    fd = os.open(path, os.O_RDONLY)
    size = os.fstat(fd).st_size
    drop_cache(fd, size)
    step = max(size // (count + 1), 4096)
    buf = 4096
    t0 = time.time()
    for i in range(1, count + 1):
        os.pread(fd, buf, i * step)
    dt = time.time() - t0
    os.close(fd)
    return {"count": count, "seconds": round(dt, 3), "iops": round(count / dt, 1) if dt else None}


def main() -> int:
    ap = argparse.ArgumentParser()
    ap.add_argument("--path", default=None, help="測試檔;沒給就用模型檔")
    ap.add_argument("--model", default=None)
    ap.add_argument("--out", required=True)
    ap.add_argument("--size-mib", type=int, default=1024)
    ap.add_argument("--ramp", type=int, default=512, help="隨機讀的次數")
    a = ap.parse_args()

    path = a.path
    tmp = None
    if not path:
        m = a.model or os.environ.get("ST_MODEL_PATH")
        if not m:
            print("需要 --path 或 --model(或 ST_MODEL_PATH)", file=sys.stderr)
            return 2
        path = m
    if not os.path.exists(path):
        tmp = f"/tmp/io-probe-{os.getpid()}.bin"
        size = a.size_mib * MIB
        print(f"[io_probe] 造 {a.size_mib} MiB 測試檔 {tmp}", file=sys.stderr)
        with open(tmp, "wb") as f:
            f.write(os.urandom(size))
        path = tmp

    size = min(a.size_mib * MIB, os.path.getsize(path))
    seq = [read_seq(path, size, t) for t in (1, 2, 4, 8) if size // t >= 8 * MIB]
    res = {
        "when": time.strftime("%Y-%m-%dT%H:%M:%S%z"),
        "file": path,
        "size_mib": size // MIB,
        "fs": os.statvfs(path).f_bsize,
        "sequential": seq,
        "random_4k": read_rand4k(path, a.ramp),
        "cpu_count": os.cpu_count(),
    }
    if tmp:
        res["tmp_removed"] = tmp
        os.unlink(tmp)
    os.makedirs(os.path.dirname(os.path.abspath(a.out)) or ".", exist_ok=True)
    with open(a.out, "w") as f:
        json.dump(res, f, indent=2)
        f.write("\n")
    print(json.dumps(res, indent=2))
    return 0


if __name__ == "__main__":
    sys.exit(main())