drkareemkamal's picture
Upload README.md with huggingface_hub
022c437 verified
|
Raw
History Blame
20.2 kB
metadata
library_name: transformers
license: mit
language:
  - en
tags:
  - medical
  - clinical-nlp
  - biobert
  - bio-clinicalbert
  - cancer
  - survival-analysis
  - oncology
  - pathology
  - tcga
  - lora
  - peft
  - cox-regression
  - risk-prediction
  - text-classification
  - feature-extraction
  - pytorch
datasets:
  - custom
base_model: emilyalsentzer/Bio_ClinicalBERT
pipeline_tag: feature-extraction
model-index:
  - name: finetunePathologicalTextUsingBioBERT
    results:
      - task:
          type: feature-extraction
          name: Survival Risk Prediction
        metrics:
          - type: loss
            name: Cox PH Validation Loss
            value: 0.529
          - type: loss
            name: Cox PH Training Loss
            value: 0.4003

🧬 Fine-Tuned Bio_ClinicalBERT for Cancer Survival Prediction from Pathological Text

A domain-adapted biomedical language model fine-tuned on 19,637 TCGA pathological text reports for cancer survival risk prediction using Cox Proportional Hazards loss with LoRA adapters β€” trained on NVIDIA RTX 3090 (24 GB VRAM).

License: MIT PyTorch Transformers PEFT GPU


Model Details

Model Description

This model is a fine-tuned version of Bio_ClinicalBERT (Alsentzer et al., 2019) adapted for cancer survival risk prediction directly from unstructured pathological text reports. The model was trained on data from The Cancer Genome Atlas (TCGA) spanning 24 cancer types across 32 cohorts.

Instead of traditional hand-crafted features (stage, grade, tumor size), this model learns survival-relevant patterns directly from raw pathological text β€” capturing subtle linguistic cues such as pathologist phrasing correlating with tumor aggressiveness, specific morphological descriptions, and diagnostic uncertainty language.

The model outputs:

  1. A continuous risk score β€” higher values indicate higher mortality risk (used with Cox Proportional Hazards framework)
  2. 768-dimensional embeddings β€” from the [CLS] token, suitable for downstream multimodal survival pipelines
  • Developed by: Dr. Kareem Kamal
  • Model type: BERT-based encoder with LoRA adapters + linear survival risk head
  • Language(s): English (clinical/biomedical)
  • License: MIT
  • Fine-tuned from: emilyalsentzer/Bio_ClinicalBERT
  • Base architecture: BERT-Base (cased, 12-layer, 768-hidden, 12-attention-heads, ~110M parameters)

Model Sources


About Bio_ClinicalBERT (Base Model)

Bio_ClinicalBERT has a unique three-stage pre-training lineage that makes it ideal for clinical text understanding:

Stage Training Data Details
1. BERT-Base Wikipedia + BookCorpus General English language understanding
2. BioBERT v1.0 PubMed abstracts (200K) + PMC full-text (270K) Biomedical scientific literature
3. Bio_ClinicalBERT MIMIC-III clinical notes (~880M words) Real electronic health records (EHR)

Key specifications of the base model:

  • Architecture: cased_L-12_H-768_A-12 (12 layers, 768 hidden dim, 12 attention heads)
  • Parameters: ~110 million
  • Vocabulary: 28,996 WordPiece tokens (domain-adapted)
  • Max sequence length: 128 tokens (original); extended to 512 tokens in our fine-tuning
  • Original training: 150,000 steps on GeForce GTX TITAN X (12 GB), batch size 32, LR 5e-5

This lineage means the model understands:

  • βœ… General English grammar and semantics (BERT)
  • βœ… Biomedical terminology and relationships (BioBERT)
  • βœ… Clinical shorthand, abbreviations, and report structure (MIMIC-III)

Uses

Direct Use

Load the fine-tuned model to extract survival-relevant embeddings or risk scores from pathological text:

from transformers import AutoTokenizer, AutoModel
import torch

# Load model and tokenizer
tokenizer = AutoTokenizer.from_pretrained("drkareemkamal/finetunePathologicalTextUsingBioBERT")
model = AutoModel.from_pretrained("drkareemkamal/finetunePathologicalTextUsingBioBERT")
model.eval()

# Example pathological report text
text = """Invasive ductal carcinoma, Nottingham grade 3/3. 
Tumor size: 2.8 cm. ER negative, PR negative, HER2 positive (3+). 
Lymphovascular invasion present. 2 of 14 sentinel lymph nodes positive 
for metastatic carcinoma. Margins: negative, closest margin 0.3 cm."""

# Tokenize
inputs = tokenizer(
    text,
    return_tensors="pt",
    max_length=512,
    truncation=True,
    padding=True
)

# Extract [CLS] embedding (768-dim)
with torch.no_grad():
    outputs = model(**inputs)
    cls_embedding = outputs.last_hidden_state[:, 0, :]  # Shape: (1, 768)

print(f"Embedding shape: {cls_embedding.shape}")  # torch.Size([1, 768])

Downstream Use

Survival Risk Scoring β€” Use with the custom risk head for direct risk prediction:

import torch.nn as nn

# Reconstruct the risk head (trained alongside the model)
risk_head = nn.Linear(768, 1)
# Load risk head weights from checkpoint if available

risk_score = risk_head(cls_embedding)
print(f"Risk score: {risk_score.item():.4f}")
# Higher score β†’ higher predicted mortality risk

Multimodal Fusion β€” Combine text embeddings with clinical, genomic, and mutation data:

# Text embedding: 768-dim from this model
# Gene expression: 50-dim from PCA of RNA-Seq FPKM values
# Mutation features: binary mutation matrix
# Clinical features: age, stage, grade, etc.

combined = torch.cat([text_emb, gene_emb, mutation_emb, clinical_emb], dim=-1)
# Feed into downstream survival model (e.g., DeepSurv, Cox-nnet)

Out-of-Scope Use

  • ❌ Not a diagnostic tool β€” This model predicts survival risk, not diagnosis
  • ❌ Not for non-cancer text β€” Trained exclusively on oncological pathology reports
  • ❌ Not for clinical deployment without regulatory approval β€” Research use only
  • ❌ Not for non-English text β€” Trained on English pathology reports only
  • ❌ Not for individual patient decisions β€” Requires human clinical oversight

Training Details

Training Data

Property Value
Source The Cancer Genome Atlas (TCGA) via cBioPortal
Dataset file merged_tcga_data_final.csv
Total samples 19,637 pathological text reports with survival outcomes
Train split 16,691 samples (85%)
Validation split 2,946 samples (15%)
Cancer types 24 disease types across 32 TCGA cohorts
Text column text β€” raw pathological report content
Survival endpoint Overall Survival: OS_MONTHS (time) + OS_STATUS (event: LIVING/DECEASED)
Event distribution ~70.7% Living / ~29.3% Deceased

Cancer type distribution in training data:

Disease Type Samples Deaths Event Rate
Adenomas and Adenocarcinomas 8,977 1,944 21.7%
Squamous Cell Neoplasms 2,764 1,166 42.2%
Ductal and Lobular Neoplasms 2,362 498 21.1%
Gliomas 1,654 794 48.0%
Cystic, Mucinous and Serous 1,078 382 35.4%
Transitional Cell Papillomas 816 386 47.3%
Others (18 types) ~1,986 varies varies

Training Procedure

Preprocessing

  1. Text cleaning: Rows with missing text, OS_MONTHS, or OS_STATUS dropped
  2. Survival labels: OS_STATUS mapped to binary events (1:DECEASED β†’ 1.0, 0:LIVING β†’ 0.0)
  3. Tokenization: WordPiece tokenizer from Bio_ClinicalBERT, max_length=512, right-truncation, max_length padding
  4. No text augmentation β€” raw pathological reports used as-is to preserve clinical accuracy

Fine-Tuning Method: LoRA (Low-Rank Adaptation)

Instead of updating all 110M parameters, we use LoRA adapters via the PEFT library to efficiently fine-tune only ~0.5% of parameters:

LoRA Parameter Value
Rank (r) 8
Alpha (Ξ±) 32
Target modules query, value (attention layers)
Dropout 0.1
Task type FEATURE_EXTRACTION
Trainable parameters 590K (0.5% of total)

Loss Function: Cox Proportional Hazards (Cox PH)

The model is trained with the negative partial log-likelihood of the Cox PH model, which:

  • Handles right-censored data (patients still alive at last follow-up)
  • Models relative hazard β€” ranking patients by risk, not predicting absolute survival time
  • Is the gold standard for survival analysis in clinical research
L(Ξ²) = -Ξ£ [log(h_i) - log(Ξ£ exp(h_j))] Γ— event_i
        i                j∈R(t_i)

Where h_i is the predicted log-hazard for patient i, and R(t_i) is the risk set at time t_i.

Training Hyperparameters

Hyperparameter Value
Optimizer AdamW
Learning rate 1e-4
Batch size 8
Max epochs 20
Early stopping patience 3 epochs
Validation split 15% (random, seed=42)
Precision FP32 (full precision)
Gradient clipping None
Scheduler None (constant LR)
Weight decay AdamW default (0.01)

Training Results

πŸ“ˆ Weights & Biases Dashboard: View Full Training Run & Loss Curves

The model was trained for all 20 epochs (early stopping was not triggered, indicating continuous improvement):

Epoch Train Loss Val Loss Best?
1 1.1658 0.9934
2 1.0408 0.9006
3 0.9440 0.8677
4 0.8720 0.8249
5 0.8122 0.7941
6 0.7347 0.7653
7 0.7011 0.7099
8 0.6649 0.7331
9 0.6167 0.6881
10 0.5849 0.6672
11 0.5562 0.6481
12 0.5424 0.6050
13 0.5150 0.6253
14 0.4998 0.6108
15 0.4705 0.5765
16 0.4630 0.6028
17 0.4347 0.5442
18 0.4230 0.5298
19 0.4104 0.5605
20 0.4003 0.5290 βœ…

Key observations:

  • Consistent downward trend in both train and validation loss over 20 epochs
  • Best validation loss: 0.5290 at epoch 20
  • Final training loss: 0.4003
  • No signs of catastrophic overfitting β€” the gap between train/val loss remains reasonable
  • Model checkpoint saved at epoch 20 (~415 MB)

Speeds, Sizes, Times

Property Value
Total training time ~4.5 hours (20 epochs on RTX 3090)
VRAM usage ~3.8 GB (FP32, batch_size=8)
Checkpoint size 415 MB (full state dict with LoRA adapters + risk head)
Embeddings output 162 MB CSV (19,637 samples Γ— 768 dimensions + risk scores)
Throughput ~120 samples/second (inference)

Evaluation

Metrics

Metric Description
Cox PH Loss Primary training objective β€” negative partial log-likelihood
C-index (Concordance Index) How well the model ranks patients by survival (0.5 = random, >0.7 = strong)
Kaplan-Meier Curves Visual separation between predicted high-risk and low-risk groups
Risk Score Distribution Separation of scores between alive vs deceased patients

Results

Metric Value
Best Validation Cox PH Loss 0.5290
Final Training Cox PH Loss 0.4003
Total epochs trained 20 / 20
Embedding dimension 768

Evaluation Outputs

The following evaluation artifacts are generated during training:

File Description
clinicalbert_training_loss.png Train vs Validation loss curves with best epoch marked
clinicalbert_training_results.csv Per-epoch numerical loss values
finetuned_text_embeddings.csv 768-dim embeddings + risk scores for all 19,637 samples

Technical Specifications

Model Architecture and Objective

Input: Raw pathological text (up to 512 tokens)
  β”‚
  β–Ό
β”Œβ”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”
β”‚     Bio_ClinicalBERT (Frozen backbone)      β”‚
β”‚     12 Transformer layers, 768 hidden dim   β”‚
β”‚     + LoRA adapters on query/value (r=8)    β”‚
β”‚     ~110M total params, ~590K trainable     β”‚
β””β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”¬β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”˜
                   β”‚
                   β–Ό
          [CLS] Token Embedding (768-dim)
                   β”‚
            β”Œβ”€β”€β”€β”€β”€β”€β”΄β”€β”€β”€β”€β”€β”€β”
            β–Ό             β–Ό
       Risk Head     Embeddings
    (Linear 768β†’1)  (768-dim vector)
            β”‚             β”‚
            β–Ό             β–Ό
     Cox PH Loss    Downstream Tasks

Compute Infrastructure

Hardware

Component Specification
GPU NVIDIA GeForce RTX 3090
GPU Memory 24,576 MiB (24 GB GDDR6X)
CUDA Compute Capability 8.6 (Ampere architecture)
NVIDIA Driver 580.126.09
CUDA Version 12.4 (PyTorch) / 13.0 (driver)

Software

Package Version
Python 3.10+
PyTorch 2.6.0+cu124
Transformers 5.7.0
PEFT 0.19.1
CUDA Toolkit 12.4
OS Linux (Ubuntu)
Package Manager uv
Experiment Tracking Weights & Biases

How to Reproduce

# 1. Clone the repository
git clone https://github.com/drkareemkamal/cancer-survival-analysis.git
cd cancer-survival-analysis

# 2. Set up environment with uv
uv venv && source .venv/bin/activate
uv sync

# 3. Configure API keys in .env
cat > .env << 'EOF'
HF_TOKEN="hf_your_huggingface_token"
HF_REPO_ID="your-username/your-repo-name"
WANDB_API_KEY="your_wandb_api_key"
WANDB_PROJECT="cancer-survival-analysis"
EOF

# 4. Run fine-tuning (baseline strategy)
python src/training/text_finetune.py

# Model will automatically push to HuggingFace Hub on completion

Fine-Tuning Strategies Available

This repository implements three fine-tuning strategies, each with both Bio_ClinicalBERT and OpenBioLLM-8B variants:

Strategy 1: Pan-Cancer Baseline (This Model)

Single model trained on all 19,637 samples. Maximum data, simplest approach.

python src/training/text_finetune.py

Strategy 2: Cancer-Type Conditioning Token

Prepends a cancer-type tag to each text to enable cancer-aware representations:

Before: "Invasive ductal carcinoma, Nottingham grade 3..."
After:  "[DUCTAL AND LOBULAR NEOPLASMS] Invasive ductal carcinoma..."
python src/training/text_finetune_conditioned.py

Strategy 3: Hierarchical Two-Stage

Stage 1 trains on all cancers, Stage 2 fine-tunes per cancer type (500+ samples):

python src/training/text_finetune_hierarchical.py

Bias, Risks, and Limitations

Dataset Bias

  • Geographic bias: TCGA data originates from US academic medical centers, which may not represent global patient populations
  • Demographic bias: The cohort reflects the demographics of TCGA participants and may underrepresent certain racial/ethnic groups
  • Institutional bias: Pathology report styles vary by institution; model performance may degrade on reports with different formatting conventions

Clinical Limitations

  • Not a diagnostic tool β€” predicts survival risk only, not disease diagnosis
  • Text quality dependency β€” performance is directly tied to report completeness and detail
  • No external validation β€” requires independent cohort validation before any clinical consideration
  • Censoring assumptions β€” Cox PH model assumes non-informative censoring, which may not always hold

Technical Limitations

  • Max 512 tokens β€” longer reports are truncated from the right, potentially losing relevant information
  • Single-modality β€” text-only; does not incorporate imaging, genomics, or structured clinical variables (see multimodal pipeline in repository)
  • FP32 only β€” not optimized for mixed-precision inference

Recommendations

  • Always pair with clinical judgment β€” this model is a decision-support tool, not a replacement for clinical expertise
  • Validate on your institution's data before use β€” report styles differ across institutions
  • Monitor for bias β€” regularly audit predictions across demographics, cancer types, and institutions
  • Regulatory compliance β€” any clinical deployment requires appropriate regulatory approval (e.g., FDA, CE marking)

Citation

BibTeX:

@software{kamal2026cancer_survival_biobert,
  title={Cancer Survival Prediction from Pathological Text Reports using Fine-Tuned Bio_ClinicalBERT},
  author={Kareem Kamal},
  year={2026},
  url={https://huggingface.co/drkareemkamal/finetunePathologicalTextUsingBioBERT},
  note={Fine-tuned on TCGA pathological reports with Cox PH loss and LoRA adapters, trained on NVIDIA RTX 3090}
}

APA:

Kamal, K. (2026). Cancer Survival Prediction from Pathological Text Reports using Fine-Tuned Bio_ClinicalBERT [Computer software]. Hugging Face. https://huggingface.co/drkareemkamal/finetunePathologicalTextUsingBioBERT


References

  1. Bio_ClinicalBERT: Alsentzer, E., et al. (2019). Publicly Available Clinical BERT Embeddings. NAACL Clinical NLP Workshop. HuggingFace | Paper
  2. BioBERT: Lee, J., et al. (2020). BioBERT: a pre-trained biomedical language representation model for biomedical text mining. Bioinformatics, 36(4), 1234–1240. Paper
  3. TCGA: The Cancer Genome Atlas Research Network. GDC Data Portal
  4. cBioPortal: Cerami, E., et al. (2012). The cBio Cancer Genomics Portal. Cancer Discovery, 2(5), 401–404. Website
  5. Cox PH Model: Cox, D.R. (1972). Regression Models and Life-Tables. Journal of the Royal Statistical Society, Series B, 34(2), 187–220.
  6. LoRA: Hu, E., et al. (2022). LoRA: Low-Rank Adaptation of Large Language Models. ICLR 2022. Paper
  7. PEFT: HuggingFace. Parameter-Efficient Fine-Tuning. GitHub

Model Card Authors

Model Card Contact


This model is for research purposes only. Always consult qualified medical professionals for clinical decisions. Not approved for clinical use.