Songline's picture
Add files using upload-large-folder tool
8999949 verified
Raw
History Blame Contribute Delete
1.66 kB
'''提供 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