'''提供 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