Image Classification
Transformers
Safetensors
Flair
vit
medical-imaging
brain-tumor
mri
vision-transformer
Instructions to use Songline/BrainTumor_FlairClassifier with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use Songline/BrainTumor_FlairClassifier with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("image-classification", model="Songline/BrainTumor_FlairClassifier") pipe("https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/hub/parrots.png")# Load model directly from transformers import AutoImageProcessor, AutoModelForImageClassification processor = AutoImageProcessor.from_pretrained("Songline/BrainTumor_FlairClassifier") model = AutoModelForImageClassification.from_pretrained("Songline/BrainTumor_FlairClassifier", device_map="auto") - Flair
How to use Songline/BrainTumor_FlairClassifier with Flair:
from flair.models import SequenceTagger tagger = SequenceTagger.load("Songline/BrainTumor_FlairClassifier") - Notebooks
- Google Colab
- Kaggle
| '''提供 FLAIR NIfTI 脑肿瘤病例级分类器''' | |
| from __future__ import annotations | |
| from dataclasses import dataclass | |
| from pathlib import Path | |
| from statistics import mean | |
| from typing import Any | |
| import nibabel as nib | |
| import numpy as np | |
| import torch | |
| from PIL import Image | |
| from transformers import AutoImageProcessor, AutoModelForImageClassification | |
| DEFAULT_THRESHOLD = 0.548381 | |
| DEFAULT_MAX_SLICES = 25 | |
| REQUIRED_MODEL_FILES = ('config.json', 'model.safetensors', 'preprocessor_config.json') | |
| class FlairClassifier: | |
| '''封装权重加载和 FLAIR NIfTI 病例级推理''' | |
| model: Any | |
| processor: Any | |
| device: torch.device | |
| model_dir: Path | |
| source: str | |
| def checkpoint_name(self) -> str: | |
| '''返回权重文件名''' | |
| return 'model.safetensors' | |
| def from_pretrained( | |
| cls, | |
| source: str | Path, | |
| *, | |
| device: str = 'auto', | |
| cache_dir: str | Path | None = None, | |
| revision: str | None = None, | |
| local_files_only: bool = False, | |
| ) -> 'FlairClassifier': | |
| '''从本地目录或 Hugging Face 模型仓库加载权重''' | |
| model_dir = _resolve_model_dir( | |
| source, | |
| cache_dir=cache_dir, | |
| revision=revision, | |
| local_files_only=local_files_only, | |
| ) | |
| missing = [name for name in REQUIRED_MODEL_FILES if not (model_dir / name).is_file()] | |
| if missing: | |
| raise FileNotFoundError(f'模型目录缺少文件 {missing}') | |
| weights_path = model_dir / 'model.safetensors' | |
| if weights_path.stat().st_size < 1_024: | |
| raise ValueError('模型权重未完整下载') | |
| runtime_device = _resolve_device(device) | |
| processor = AutoImageProcessor.from_pretrained( | |
| model_dir, | |
| local_files_only=True, | |
| use_fast=False, | |
| ) | |
| model = AutoModelForImageClassification.from_pretrained( | |
| model_dir, | |
| local_files_only=True, | |
| use_safetensors=True, | |
| ).to(runtime_device) | |
| model.eval() | |
| labels = {str(name).strip().lower(): int(index) for name, index in model.config.label2id.items()} | |
| if labels != {'no': 0, 'yes': 1}: | |
| raise ValueError(f'模型标签不符合预期 {labels}') | |
| return cls( | |
| model=model, | |
| processor=processor, | |
| device=runtime_device, | |
| model_dir=model_dir, | |
| source=str(source), | |
| ) | |
| def predict_nifti( | |
| self, | |
| path: str | Path, | |
| *, | |
| threshold: float = DEFAULT_THRESHOLD, | |
| max_slices: int = DEFAULT_MAX_SLICES, | |
| batch_size: int = 25, | |
| ) -> dict[str, Any]: | |
| '''对单份 FLAIR NIfTI 返回病例级分类与切片证据''' | |
| if not 0 < threshold < 1: | |
| raise ValueError('threshold 必须位于 0 到 1 之间') | |
| if batch_size <= 0: | |
| raise ValueError('batch_size 必须大于 0') | |
| volume_path = Path(path).expanduser().resolve() | |
| volume, indices, canonical_shape, foreground_slices, intensity_window = _prepare_volume( | |
| volume_path, | |
| max_slices=max_slices, | |
| ) | |
| scores = self._predict_slices(volume, batch_size=batch_size) | |
| yes_probability = mean(scores) | |
| predicted_class = 'yes' if yes_probability >= threshold else 'no' | |
| return { | |
| 'source': self.source, | |
| 'model_dir': str(self.model_dir), | |
| 'checkpoint': 'model.safetensors', | |
| 'device': str(self.device), | |
| 'input_file': str(volume_path), | |
| 'threshold': round(threshold, 6), | |
| 'predicted_class': predicted_class, | |
| 'yes_probability': round(yes_probability, 6), | |
| 'no_probability': round(1 - yes_probability, 6), | |
| 'evaluated_slices': len(scores), | |
| 'slice_indices': indices, | |
| 'slice_yes_probabilities': [round(score, 6) for score in scores], | |
| 'canonical_shape': list(canonical_shape), | |
| 'foreground_slices': foreground_slices, | |
| 'intensity_window': list(intensity_window), | |
| } | |
| def predict_images(self, images: list[Image.Image], *, batch_size: int) -> list[dict[str, Any]]: | |
| '''返回与平台兼容的逐切片分类结果''' | |
| if not images: | |
| raise ValueError('至少需要一张切片') | |
| if batch_size <= 0: | |
| raise ValueError('batch_size 必须大于 0') | |
| scores = self._predict_slices(images, batch_size=batch_size) | |
| return [ | |
| { | |
| 'class': 'yes' if score >= 0.5 else 'no', | |
| 'class_id': 1 if score >= 0.5 else 0, | |
| 'confidence': max(score, 1 - score), | |
| 'probabilities': {'no': 1 - score, 'yes': score}, | |
| } | |
| for score in scores | |
| ] | |
| def _predict_slices(self, images: list[Image.Image], *, batch_size: int) -> list[float]: | |
| '''计算每张切片的阳性概率''' | |
| scores: list[float] = [] | |
| with torch.inference_mode(): | |
| for start in range(0, len(images), batch_size): | |
| inputs = self.processor(images=images[start : start + batch_size], return_tensors='pt') | |
| inputs = {key: value.to(self.device) for key, value in inputs.items()} | |
| probabilities = torch.softmax(self.model(**inputs).logits, dim=-1).cpu().numpy() | |
| scores.extend(float(row[1]) for row in probabilities) | |
| return scores | |
| def predict_images( | |
| classifier: FlairClassifier, | |
| images: list[Image.Image], | |
| *, | |
| batch_size: int, | |
| ) -> list[dict[str, Any]]: | |
| '''以函数形式提供平台兼容的逐切片推理接口''' | |
| return classifier.predict_images(images, batch_size=batch_size) | |
| def _resolve_model_dir( | |
| source: str | Path, | |
| *, | |
| cache_dir: str | Path | None, | |
| revision: str | None, | |
| local_files_only: bool, | |
| ) -> Path: | |
| '''解析本地目录或下载模型仓库快照''' | |
| local_path = Path(source).expanduser() | |
| if local_path.is_dir(): | |
| return local_path.resolve() | |
| try: | |
| from huggingface_hub import snapshot_download | |
| except ImportError as exc: | |
| raise RuntimeError('远程模型加载需要安装 huggingface-hub') from exc | |
| return Path( | |
| snapshot_download( | |
| repo_id=str(source), | |
| revision=revision, | |
| cache_dir=str(cache_dir) if cache_dir else None, | |
| local_files_only=local_files_only, | |
| allow_patterns=list(REQUIRED_MODEL_FILES), | |
| ) | |
| ).resolve() | |
| def _resolve_device(name: str) -> torch.device: | |
| '''解析推理设备''' | |
| if name == 'auto': | |
| return torch.device('cuda' if torch.cuda.is_available() else 'cpu') | |
| runtime_device = torch.device(name) | |
| if runtime_device.type == 'cuda' and not torch.cuda.is_available(): | |
| raise RuntimeError('当前环境没有可用 CUDA 设备') | |
| return runtime_device | |
| def _prepare_volume( | |
| path: Path, | |
| *, | |
| max_slices: int, | |
| ) -> tuple[list[Image.Image], list[int], tuple[int, int, int], int, tuple[float, float]]: | |
| '''读取 FLAIR 并均匀采样有效轴位切片''' | |
| if max_slices <= 0: | |
| raise ValueError('max_slices 必须大于 0') | |
| try: | |
| data = np.asarray(nib.as_closest_canonical(nib.load(str(path))).dataobj, dtype=np.float32) | |
| except Exception as exc: | |
| raise ValueError(f'无法读取输入文件 {path}') from exc | |
| if data.ndim != 3: | |
| raise ValueError(f'输入必须是三维 NIfTI 文件 当前形状为 {data.shape}') | |
| if not np.isfinite(data).all(): | |
| raise ValueError('输入包含 NaN 或无穷值') | |
| foreground = data != 0 | |
| if not foreground.any(): | |
| raise ValueError('输入体积为空') | |
| counts = foreground.sum(axis=(0, 1)) | |
| candidates = np.flatnonzero(counts >= max(1, int(counts.max() * 0.01))) | |
| positions = np.linspace(0, candidates.size - 1, num=min(max_slices, candidates.size)) | |
| indices = candidates[np.rint(positions).astype(np.int64)] | |
| lower, upper = np.percentile(data[foreground].astype(np.float64, copy=False), (1, 99)) | |
| if not np.isfinite(lower) or not np.isfinite(upper) or upper <= lower: | |
| raise ValueError('输入没有有效的强度变化') | |
| images = [] | |
| for index in indices: | |
| normalized = np.clip((data[:, :, int(index)] - lower) / (upper - lower), 0, 1) | |
| pixels = np.rint(normalized * 255).astype(np.uint8) | |
| images.append(Image.fromarray(pixels, mode='L').convert('RGB')) | |
| return ( | |
| images, | |
| [int(index) for index in indices], | |
| tuple(int(size) for size in data.shape), | |
| int(candidates.size), | |
| (round(float(lower), 6), round(float(upper), 6)), | |
| ) | |