File size: 1,655 Bytes
8999949
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
'''提供 FLAIR 分类命令行入口'''

from __future__ import annotations

import argparse
import json
from pathlib import Path

from .classifier import DEFAULT_MAX_SLICES, DEFAULT_THRESHOLD, FlairClassifier


def parse_args() -> argparse.Namespace:
    '''解析命令行参数'''

    parser = argparse.ArgumentParser(description='FLAIR NIfTI 脑肿瘤二分类')
    parser.add_argument('--input', required=True, type=Path, help='FLAIR NIfTI 文件路径')
    parser.add_argument('--source', default=Path(__file__).resolve().parents[2], help='本地模型目录或 Hugging Face 模型 ID')
    parser.add_argument('--output', type=Path, help='JSON 输出文件路径')
    parser.add_argument('--threshold', type=float, default=DEFAULT_THRESHOLD, help='阳性判定阈值')
    parser.add_argument('--max-slices', type=int, default=DEFAULT_MAX_SLICES, help='最大采样切片数')
    parser.add_argument('--batch-size', type=int, default=25, help='推理批次大小')
    parser.add_argument('--device', default='auto', help='推理设备')
    return parser.parse_args()


def main() -> int:
    '''执行单病例分类'''

    args = parse_args()
    classifier = FlairClassifier.from_pretrained(args.source, device=args.device)
    result = classifier.predict_nifti(
        args.input,
        threshold=args.threshold,
        max_slices=args.max_slices,
        batch_size=args.batch_size,
    )
    content = json.dumps(result, ensure_ascii=False, indent=2)
    if args.output:
        args.output.parent.mkdir(parents=True, exist_ok=True)
        args.output.write_text(content + '\n', encoding='utf-8')
    print(content)
    return 0