Repository files navigation

TLDChoiceNet: Quantitatively Choosing a Transfer Learning Dataset

Predict how well a transfer-learning dataset will work — before spending the compute to fine-tune on it.

PosterPython 3.10TensorFlowCode style: black

Jing Ning, James D. Braza — Stanford University ·Paper (https://arxiv.org/abs/2608.09091)

Project poster

The original CS 330 poster — click for the full-resolution PDF. It covers TLDChoiceNet v1; the v2 architecture and the ACC metric below came later, so the figures in the poster differ from the results reported here.

Joint work. Originally developed at jamesbraza/cs330-project; this repository preserves the full commit history.


The problem

You have a small dataset for a target task. You have candidate pre-trained datasets A, B, and C — similar example counts, similar class counts — but compute to fine-tune on only one of them. Which do you pick?

Conventional wisdom says take the biggest, most diverse source. But no quantitative method existed to make that call. TLDChoiceNet predicts the post-fine-tuning test accuracy for each candidate, so the choice becomes a measurement instead of a guess.

Results

Model / metricResult
TLDChoiceNet v1 — test MSE0.154
TLDChoiceNet v2 — test MSE0.031 (5× lower)
Distribution distance (DD) vs. accuracy — 0.894
Average class correlation (ACC) vs. accuracy — 0.974

Two results worth separating:

The learned predictor. v2 cuts test MSE 5× over v1 by embedding both inputs per class rather than averaging across all of them, and by learning the post-embedding reduction rather than fixing it. It has ~1.15M trainable parameters against v1's ~314K.

The unsupervised metrics — no training required at all. Average class correlation reaches an of 0.974 against fine-tune accuracy. You can compute it from a pre-trained ResNet50 v2 forward pass and pick your transfer dataset without training anything.

Approach

The Transfer Learning Dataset (TLDS)

The core obstacle is that there is no dataset of "transfer learning dataset → resulting accuracy" pairs to learn from. So we built one. Each entry is a 3-tuple of (transfer learning dataset, transfer-learned model, resulting test accuracy), spanning four deliberately chosen regimes:

RegimeSourceSubsets
SimilarPlant diseases3 × 10 classes
DissimilarBird species3 × 10 classes
RandomCIFAR-100 and ImageNet40 × 10 classes
No transfer learningRandom initialisation (control)1

Each invocation of the generation script produces 57 data points; varying the seed yields more. Fine-tuning throughout is on a 22-class plant-leaves dataset.

TLDChoiceNet

Both versions take an embedded fine-tuning dataset and an embedded transfer-learning dataset/model, and regress to a single number — predicted test accuracy.

  • v1 embeds the fine-tuning dataset via PCA to 256 dimensions and the transfer-learned model via its flattened last Conv2D, applies a LoRA-similar learned reduction, and concatenates.
  • v2 fixes v1's central flaw: both embeddings ignored class specifics. v2 embeds per class, keeps the 10 highest-activation classes, and adds the two reduced matrices instead of concatenating them.

Two unsupervised metrics

Distribution distance (DD) — the Euclidean distance between two datasets in (mean, |skew|, |kurtosis|) space over normalised pixel values:

$$\mathrm{DD}(i,j)=\sqrt{(\mu_i-\mu_j)^2+(s_i-s_j)^2+(\kappa_i-\kappa_j)^2}$$

Average class correlation (ACC) — mean pairwise correlation between per-class embeddings from an ImageNet-pretrained headless ResNet50 v2:

$$\mathrm{ACC}(i,j)=\frac{1}{nm}\sum_{k=1}^{n}\sum_{l=1}^{m}\mathrm{Cor}_{kl}$$

ACC separates similar from dissimilar transfer datasets far more sharply (0.501 vs. 0.276 — 45% lower for dissimilar) and explains fine-tune accuracy better than DD does.

Key findings

Pre-trained weights push dissimilar classes apart in latent space. Computing ACC from raw normalised pixels gives 0.507 for similar classes and 0.438 for dissimilar — a gap of just 0.06. Running the same computation through ImageNet-pretrained ResNet50 v2 weights widens that gap to 0.22 (0.49 vs. 0.27). The pre-trained weights are actively increasing class separation, which is a concrete answer to what transfer learning transfers.

A dataset's low-level pixel statistics explain much of the transfer effect. DD, computed from nothing but mean, skew, and kurtosis of pixel values, reaches an of 0.894.

Similar transfer datasets move weights less during fine-tuning. L2 distance between pre-trained and fine-tuned weights stays smaller across epochs for similar transfer datasets, corroborated by centered kernel alignment (CKA).

Honest limitations

  • The reported MSE figures come from a test subset that shares its fine-tuning dataset with the training subset. On a genuinely unseen fine-tuning dataset, v1 degrades badly — MSE 0.46, with accuracy predictions off by as much as 80%. Generalising across fine-tuning datasets needs either a v3 architecture or a TLDS spanning multiple fine-tuning datasets.
  • Test accuracy spans only ~15% from random initialisation to the best transfer-learned model, which is a narrow band in which to resolve differences.
  • Dissimilar transfer datasets scored about the same as random ones, so image diversity alone may matter less than expected here.
  • Everything is image classification on one fine-tuning task. Generalisation to other domains is untested.

Datasets

DatasetRole
Plant leaves (22 classes)Fine-tuning target
Plant diseases / plant_villageSimilar transfer source
Birds 450 speciesDissimilar transfer source
CIFAR-100, ImageNetRandom transfer sources

Download them all with the Kaggle API:

kaggle datasets download -p data/plant-diseases --unzip vipoooool/new-plant-diseases-dataset
kaggle datasets download -p data/plant-leaves --unzip csafrit2/plant-leaves-for-image-classification
kaggle datasets download -p data/bird-species --unzip gpiosenka/100-bird-species

Getting started

Developed with Python 3.10.

python -m venv venv
source venv/bin/activate
python -m pip install -r requirements.txt

The pipeline runs in three stages — build the TLDS, fine-tune, then train ChoiceNet on the result:

bash training/1.run_tl_training.sh # train TransferModel on each TL dataset
bash training/2.run_fine_tune.sh # fine-tune onto plant leaves, record accuracy
bash training/3.run_choicenet.sh # train TLDChoiceNet on the resulting TLDS

The shell scripts carry absolute paths from the original GPU machine (/data1/cs330/project/...). Point them at your own directories before running.

Monitor training:

tensorboard --logdir training # then open http://localhost:6006/

Code quality tooling

python -m pip install -r requirements-qa.txt
pre-commit install

Repository layout

data/ Dataset loading, preprocessing, and TLDS source configs
models/ TransferModel CNN and ChoiceNet v1/v2 architectures
training/ TLDS creation, fine-tuning, and ChoiceNet training entry points
embedding/ ResNet50 v2 dataset embedding and weight-matrix preprocessing
experiments/ Metric analysis — CKA, KS tests, correlation matrices, weight distance

Citation

@article{ning2026tldchoicenet,
title = {TLDChoiceNet: Quantitatively Choosing a Transfer Learning Dataset},
author = {Ning, Jing and Braza, James D.},
journal = {arXiv preprint},
year = {2026}
}

Acknowledgements

We thank Chelsea Finn and Daniel Zeng for helpful discussions and feedback, including ideas on ImageNet embedding and measuring weight distance across training.

About

Predict how well a transfer-learning dataset will work before fine-tuning on it. TLDChoiceNet cuts prediction MSE 5x, and an unsupervised class-correlation metric explains fine-tune accuracy with R2=0.97. Stanford CS 330.

Topics

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages

, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Add copy buttons to all
 blocks\n(function() {\n function addCopyButtons() {\n document.querySelectorAll('pre code').forEach(function(codeBlock) {\n if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;\n codeBlock.parentElement.setAttribute('data-copy-added', 'true');\n \n var btn = document.createElement('button');\n btn.textContent = 'Copy';\n btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';\n btn.onmouseover = function() { this.style.opacity = '1'; };\n btn.onmouseout = function() { this.style.opacity = '0.7'; };\n btn.onclick = function() {\n navigator.clipboard.writeText(codeBlock.textContent).then(function() {\n btn.textContent = 'Copied!';\n setTimeout(function() { btn.textContent = 'Copy'; }, 1500);\n });\n };\n codeBlock.parentElement.style.position = 'relative';\n codeBlock.parentElement.appendChild(btn);\n });\n }\n \n addCopyButtons();\n \n // Re-run on dynamic content\n var observer = new MutationObserver(addCopyButtons);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Add Copy Buttons to Code Blocks");
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
Skip to content

Repository files navigation

TLDChoiceNet: Quantitatively Choosing a Transfer Learning Dataset

Predict how well a transfer-learning dataset will work — before spending the compute to fine-tune on it.

PosterPython 3.10TensorFlowCode style: black

Jing Ning, James D. Braza — Stanford University ·Paper (https://arxiv.org/abs/2608.09091)

Project poster

The original CS 330 poster — click for the full-resolution PDF. It covers TLDChoiceNet v1; the v2 architecture and the ACC metric below came later, so the figures in the poster differ from the results reported here.

Joint work. Originally developed at jamesbraza/cs330-project; this repository preserves the full commit history.


The problem

You have a small dataset for a target task. You have candidate pre-trained datasets A, B, and C — similar example counts, similar class counts — but compute to fine-tune on only one of them. Which do you pick?

Conventional wisdom says take the biggest, most diverse source. But no quantitative method existed to make that call. TLDChoiceNet predicts the post-fine-tuning test accuracy for each candidate, so the choice becomes a measurement instead of a guess.

Results

Model / metricResult
TLDChoiceNet v1 — test MSE0.154
TLDChoiceNet v2 — test MSE0.031 (5× lower)
Distribution distance (DD) vs. accuracy — 0.894
Average class correlation (ACC) vs. accuracy — 0.974

Two results worth separating:

The learned predictor. v2 cuts test MSE 5× over v1 by embedding both inputs per class rather than averaging across all of them, and by learning the post-embedding reduction rather than fixing it. It has ~1.15M trainable parameters against v1's ~314K.

The unsupervised metrics — no training required at all. Average class correlation reaches an of 0.974 against fine-tune accuracy. You can compute it from a pre-trained ResNet50 v2 forward pass and pick your transfer dataset without training anything.

Approach

The Transfer Learning Dataset (TLDS)

The core obstacle is that there is no dataset of "transfer learning dataset → resulting accuracy" pairs to learn from. So we built one. Each entry is a 3-tuple of (transfer learning dataset, transfer-learned model, resulting test accuracy), spanning four deliberately chosen regimes:

RegimeSourceSubsets
SimilarPlant diseases3 × 10 classes
DissimilarBird species3 × 10 classes
RandomCIFAR-100 and ImageNet40 × 10 classes
No transfer learningRandom initialisation (control)1

Each invocation of the generation script produces 57 data points; varying the seed yields more. Fine-tuning throughout is on a 22-class plant-leaves dataset.

TLDChoiceNet

Both versions take an embedded fine-tuning dataset and an embedded transfer-learning dataset/model, and regress to a single number — predicted test accuracy.

  • v1 embeds the fine-tuning dataset via PCA to 256 dimensions and the transfer-learned model via its flattened last Conv2D, applies a LoRA-similar learned reduction, and concatenates.
  • v2 fixes v1's central flaw: both embeddings ignored class specifics. v2 embeds per class, keeps the 10 highest-activation classes, and adds the two reduced matrices instead of concatenating them.

Two unsupervised metrics

Distribution distance (DD) — the Euclidean distance between two datasets in (mean, |skew|, |kurtosis|) space over normalised pixel values:

$$\mathrm{DD}(i,j)=\sqrt{(\mu_i-\mu_j)^2+(s_i-s_j)^2+(\kappa_i-\kappa_j)^2}$$

Average class correlation (ACC) — mean pairwise correlation between per-class embeddings from an ImageNet-pretrained headless ResNet50 v2:

$$\mathrm{ACC}(i,j)=\frac{1}{nm}\sum_{k=1}^{n}\sum_{l=1}^{m}\mathrm{Cor}_{kl}$$

ACC separates similar from dissimilar transfer datasets far more sharply (0.501 vs. 0.276 — 45% lower for dissimilar) and explains fine-tune accuracy better than DD does.

Key findings

Pre-trained weights push dissimilar classes apart in latent space. Computing ACC from raw normalised pixels gives 0.507 for similar classes and 0.438 for dissimilar — a gap of just 0.06. Running the same computation through ImageNet-pretrained ResNet50 v2 weights widens that gap to 0.22 (0.49 vs. 0.27). The pre-trained weights are actively increasing class separation, which is a concrete answer to what transfer learning transfers.

A dataset's low-level pixel statistics explain much of the transfer effect. DD, computed from nothing but mean, skew, and kurtosis of pixel values, reaches an of 0.894.

Similar transfer datasets move weights less during fine-tuning. L2 distance between pre-trained and fine-tuned weights stays smaller across epochs for similar transfer datasets, corroborated by centered kernel alignment (CKA).

Honest limitations

  • The reported MSE figures come from a test subset that shares its fine-tuning dataset with the training subset. On a genuinely unseen fine-tuning dataset, v1 degrades badly — MSE 0.46, with accuracy predictions off by as much as 80%. Generalising across fine-tuning datasets needs either a v3 architecture or a TLDS spanning multiple fine-tuning datasets.
  • Test accuracy spans only ~15% from random initialisation to the best transfer-learned model, which is a narrow band in which to resolve differences.
  • Dissimilar transfer datasets scored about the same as random ones, so image diversity alone may matter less than expected here.
  • Everything is image classification on one fine-tuning task. Generalisation to other domains is untested.

Datasets

DatasetRole
Plant leaves (22 classes)Fine-tuning target
Plant diseases / plant_villageSimilar transfer source
Birds 450 speciesDissimilar transfer source
CIFAR-100, ImageNetRandom transfer sources

Download them all with the Kaggle API:

kaggle datasets download -p data/plant-diseases --unzip vipoooool/new-plant-diseases-dataset
kaggle datasets download -p data/plant-leaves --unzip csafrit2/plant-leaves-for-image-classification
kaggle datasets download -p data/bird-species --unzip gpiosenka/100-bird-species

Getting started

Developed with Python 3.10.

python -m venv venv
source venv/bin/activate
python -m pip install -r requirements.txt

The pipeline runs in three stages — build the TLDS, fine-tune, then train ChoiceNet on the result:

bash training/1.run_tl_training.sh # train TransferModel on each TL dataset
bash training/2.run_fine_tune.sh # fine-tune onto plant leaves, record accuracy
bash training/3.run_choicenet.sh # train TLDChoiceNet on the resulting TLDS

The shell scripts carry absolute paths from the original GPU machine (/data1/cs330/project/...). Point them at your own directories before running.

Monitor training:

tensorboard --logdir training # then open http://localhost:6006/

Code quality tooling

python -m pip install -r requirements-qa.txt
pre-commit install

Repository layout

data/ Dataset loading, preprocessing, and TLDS source configs
models/ TransferModel CNN and ChoiceNet v1/v2 architectures
training/ TLDS creation, fine-tuning, and ChoiceNet training entry points
embedding/ ResNet50 v2 dataset embedding and weight-matrix preprocessing
experiments/ Metric analysis — CKA, KS tests, correlation matrices, weight distance

Citation

@article{ning2026tldchoicenet,
title = {TLDChoiceNet: Quantitatively Choosing a Transfer Learning Dataset},
author = {Ning, Jing and Braza, James D.},
journal = {arXiv preprint},
year = {2026}
}

Acknowledgements

We thank Chelsea Finn and Daniel Zeng for helpful discussions and feedback, including ideas on ImageNet embedding and measuring weight distance across training.

About

Predict how well a transfer-learning dataset will work before fine-tuning on it. TLDChoiceNet cuts prediction MSE 5x, and an unsupervised class-correlation metric explains fine-tune accuracy with R2=0.97. Stanford CS 330.

Topics

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages

, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Force GitHub README to respect dark mode\n(function() {\n var style = document.createElement('style');\n style.textContent = '\n .markdown-body {\n color-scheme: dark light;\n }\n .markdown-body pre { background: #161b22 !important; }\n .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; }\n .markdown-body table th, .markdown-body table td { border-color: #30363d !important; }\n .markdown-body img { background: #0d1117; }\n .markdown-body blockquote { border-left-color: #8b949e; }\n .markdown-body hr { border-color: #30363d; }\n ';\n document.head.appendChild(style);\n})();", "GitHub Dark Mode README Fix"); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content

Repository files navigation

TLDChoiceNet: Quantitatively Choosing a Transfer Learning Dataset

Predict how well a transfer-learning dataset will work — before spending the compute to fine-tune on it.

PosterPython 3.10TensorFlowCode style: black

Jing Ning, James D. Braza — Stanford University ·Paper (https://arxiv.org/abs/2608.09091)

Project poster

The original CS 330 poster — click for the full-resolution PDF. It covers TLDChoiceNet v1; the v2 architecture and the ACC metric below came later, so the figures in the poster differ from the results reported here.

Joint work. Originally developed at jamesbraza/cs330-project; this repository preserves the full commit history.


The problem

You have a small dataset for a target task. You have candidate pre-trained datasets A, B, and C — similar example counts, similar class counts — but compute to fine-tune on only one of them. Which do you pick?

Conventional wisdom says take the biggest, most diverse source. But no quantitative method existed to make that call. TLDChoiceNet predicts the post-fine-tuning test accuracy for each candidate, so the choice becomes a measurement instead of a guess.

Results

Model / metricResult
TLDChoiceNet v1 — test MSE0.154
TLDChoiceNet v2 — test MSE0.031 (5× lower)
Distribution distance (DD) vs. accuracy — 0.894
Average class correlation (ACC) vs. accuracy — 0.974

Two results worth separating:

The learned predictor. v2 cuts test MSE 5× over v1 by embedding both inputs per class rather than averaging across all of them, and by learning the post-embedding reduction rather than fixing it. It has ~1.15M trainable parameters against v1's ~314K.

The unsupervised metrics — no training required at all. Average class correlation reaches an of 0.974 against fine-tune accuracy. You can compute it from a pre-trained ResNet50 v2 forward pass and pick your transfer dataset without training anything.

Approach

The Transfer Learning Dataset (TLDS)

The core obstacle is that there is no dataset of "transfer learning dataset → resulting accuracy" pairs to learn from. So we built one. Each entry is a 3-tuple of (transfer learning dataset, transfer-learned model, resulting test accuracy), spanning four deliberately chosen regimes:

RegimeSourceSubsets
SimilarPlant diseases3 × 10 classes
DissimilarBird species3 × 10 classes
RandomCIFAR-100 and ImageNet40 × 10 classes
No transfer learningRandom initialisation (control)1

Each invocation of the generation script produces 57 data points; varying the seed yields more. Fine-tuning throughout is on a 22-class plant-leaves dataset.

TLDChoiceNet

Both versions take an embedded fine-tuning dataset and an embedded transfer-learning dataset/model, and regress to a single number — predicted test accuracy.

  • v1 embeds the fine-tuning dataset via PCA to 256 dimensions and the transfer-learned model via its flattened last Conv2D, applies a LoRA-similar learned reduction, and concatenates.
  • v2 fixes v1's central flaw: both embeddings ignored class specifics. v2 embeds per class, keeps the 10 highest-activation classes, and adds the two reduced matrices instead of concatenating them.

Two unsupervised metrics

Distribution distance (DD) — the Euclidean distance between two datasets in (mean, |skew|, |kurtosis|) space over normalised pixel values:

$$\mathrm{DD}(i,j)=\sqrt{(\mu_i-\mu_j)^2+(s_i-s_j)^2+(\kappa_i-\kappa_j)^2}$$

Average class correlation (ACC) — mean pairwise correlation between per-class embeddings from an ImageNet-pretrained headless ResNet50 v2:

$$\mathrm{ACC}(i,j)=\frac{1}{nm}\sum_{k=1}^{n}\sum_{l=1}^{m}\mathrm{Cor}_{kl}$$

ACC separates similar from dissimilar transfer datasets far more sharply (0.501 vs. 0.276 — 45% lower for dissimilar) and explains fine-tune accuracy better than DD does.

Key findings

Pre-trained weights push dissimilar classes apart in latent space. Computing ACC from raw normalised pixels gives 0.507 for similar classes and 0.438 for dissimilar — a gap of just 0.06. Running the same computation through ImageNet-pretrained ResNet50 v2 weights widens that gap to 0.22 (0.49 vs. 0.27). The pre-trained weights are actively increasing class separation, which is a concrete answer to what transfer learning transfers.

A dataset's low-level pixel statistics explain much of the transfer effect. DD, computed from nothing but mean, skew, and kurtosis of pixel values, reaches an of 0.894.

Similar transfer datasets move weights less during fine-tuning. L2 distance between pre-trained and fine-tuned weights stays smaller across epochs for similar transfer datasets, corroborated by centered kernel alignment (CKA).

Honest limitations

  • The reported MSE figures come from a test subset that shares its fine-tuning dataset with the training subset. On a genuinely unseen fine-tuning dataset, v1 degrades badly — MSE 0.46, with accuracy predictions off by as much as 80%. Generalising across fine-tuning datasets needs either a v3 architecture or a TLDS spanning multiple fine-tuning datasets.
  • Test accuracy spans only ~15% from random initialisation to the best transfer-learned model, which is a narrow band in which to resolve differences.
  • Dissimilar transfer datasets scored about the same as random ones, so image diversity alone may matter less than expected here.
  • Everything is image classification on one fine-tuning task. Generalisation to other domains is untested.

Datasets

DatasetRole
Plant leaves (22 classes)Fine-tuning target
Plant diseases / plant_villageSimilar transfer source
Birds 450 speciesDissimilar transfer source
CIFAR-100, ImageNetRandom transfer sources

Download them all with the Kaggle API:

kaggle datasets download -p data/plant-diseases --unzip vipoooool/new-plant-diseases-dataset
kaggle datasets download -p data/plant-leaves --unzip csafrit2/plant-leaves-for-image-classification
kaggle datasets download -p data/bird-species --unzip gpiosenka/100-bird-species

Getting started

Developed with Python 3.10.

python -m venv venv
source venv/bin/activate
python -m pip install -r requirements.txt

The pipeline runs in three stages — build the TLDS, fine-tune, then train ChoiceNet on the result:

bash training/1.run_tl_training.sh # train TransferModel on each TL dataset
bash training/2.run_fine_tune.sh # fine-tune onto plant leaves, record accuracy
bash training/3.run_choicenet.sh # train TLDChoiceNet on the resulting TLDS

The shell scripts carry absolute paths from the original GPU machine (/data1/cs330/project/...). Point them at your own directories before running.

Monitor training:

tensorboard --logdir training # then open http://localhost:6006/

Code quality tooling

python -m pip install -r requirements-qa.txt
pre-commit install

Repository layout

data/ Dataset loading, preprocessing, and TLDS source configs
models/ TransferModel CNN and ChoiceNet v1/v2 architectures
training/ TLDS creation, fine-tuning, and ChoiceNet training entry points
embedding/ ResNet50 v2 dataset embedding and weight-matrix preprocessing
experiments/ Metric analysis — CKA, KS tests, correlation matrices, weight distance

Citation

@article{ning2026tldchoicenet,
title = {TLDChoiceNet: Quantitatively Choosing a Transfer Learning Dataset},
author = {Ning, Jing and Braza, James D.},
journal = {arXiv preprint},
year = {2026}
}

Acknowledgements

We thank Chelsea Finn and Daniel Zeng for helpful discussions and feedback, including ideas on ImageNet embedding and measuring weight distance across training.

About

Predict how well a transfer-learning dataset will work before fine-tuning on it. TLDChoiceNet cuts prediction MSE 5x, and an unsupervised class-correlation metric explains fine-tune accuracy with R2=0.97. Stanford CS 330.

Topics

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages

, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Highlight search terms from Google/DuckDuckGo/Bing referrer\n(function() {\n var ref = document.referrer;\n var terms = [];\n \n if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) {\n var url = new URL(ref);\n var q = url.searchParams.get('q') || url.searchParams.get('p');\n if (q) {\n terms = q.split(/\\s+/).filter(function(t) { return t.length > 2; });\n }\n }\n \n if (terms.length === 0) return;\n \n var style = document.createElement('style');\n style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }';\n document.head.appendChild(style);\n \n function highlight(node) {\n if (node.nodeType === 3) { // text node\n var text = node.textContent;\n var found = false;\n terms.forEach(function(term) {\n var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\\]\\\\]/g, '\\\\') + ')', 'gi');\n if (regex.test(text)) {\n found = true;\n var frag = document.createDocumentFragment();\n var parts = text.split(regex);\n parts.forEach(function(part, i) {\n if (i % 2 === 0) {\n frag.appendChild(document.createTextNode(part));\n } else {\n var span = document.createElement('span');\n span.className = 'userscript-highlight';\n span.textContent = part;\n frag.appendChild(span);\n }\n });\n node.parentNode.replaceChild(frag, node);\n }\n });\n } else if (node.nodeType === 1 && node.childNodes) { // element\n var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT'];\n if (!skipTags.includes(node.tagName)) {\n Array.from(node.childNodes).forEach(highlight);\n }\n }\n }\n \n highlight(document.body);\n \n // Re-highlight on dynamic content\n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1 || node.nodeType === 3) highlight(node);\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Highlight Search Terms"); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content

Repository files navigation

TLDChoiceNet: Quantitatively Choosing a Transfer Learning Dataset

Predict how well a transfer-learning dataset will work — before spending the compute to fine-tune on it.

PosterPython 3.10TensorFlowCode style: black

Jing Ning, James D. Braza — Stanford University ·Paper (https://arxiv.org/abs/2608.09091)

Project poster

The original CS 330 poster — click for the full-resolution PDF. It covers TLDChoiceNet v1; the v2 architecture and the ACC metric below came later, so the figures in the poster differ from the results reported here.

Joint work. Originally developed at jamesbraza/cs330-project; this repository preserves the full commit history.


The problem

You have a small dataset for a target task. You have candidate pre-trained datasets A, B, and C — similar example counts, similar class counts — but compute to fine-tune on only one of them. Which do you pick?

Conventional wisdom says take the biggest, most diverse source. But no quantitative method existed to make that call. TLDChoiceNet predicts the post-fine-tuning test accuracy for each candidate, so the choice becomes a measurement instead of a guess.

Results

Model / metricResult
TLDChoiceNet v1 — test MSE0.154
TLDChoiceNet v2 — test MSE0.031 (5× lower)
Distribution distance (DD) vs. accuracy — 0.894
Average class correlation (ACC) vs. accuracy — 0.974

Two results worth separating:

The learned predictor. v2 cuts test MSE 5× over v1 by embedding both inputs per class rather than averaging across all of them, and by learning the post-embedding reduction rather than fixing it. It has ~1.15M trainable parameters against v1's ~314K.

The unsupervised metrics — no training required at all. Average class correlation reaches an of 0.974 against fine-tune accuracy. You can compute it from a pre-trained ResNet50 v2 forward pass and pick your transfer dataset without training anything.

Approach

The Transfer Learning Dataset (TLDS)

The core obstacle is that there is no dataset of "transfer learning dataset → resulting accuracy" pairs to learn from. So we built one. Each entry is a 3-tuple of (transfer learning dataset, transfer-learned model, resulting test accuracy), spanning four deliberately chosen regimes:

RegimeSourceSubsets
SimilarPlant diseases3 × 10 classes
DissimilarBird species3 × 10 classes
RandomCIFAR-100 and ImageNet40 × 10 classes
No transfer learningRandom initialisation (control)1

Each invocation of the generation script produces 57 data points; varying the seed yields more. Fine-tuning throughout is on a 22-class plant-leaves dataset.

TLDChoiceNet

Both versions take an embedded fine-tuning dataset and an embedded transfer-learning dataset/model, and regress to a single number — predicted test accuracy.

  • v1 embeds the fine-tuning dataset via PCA to 256 dimensions and the transfer-learned model via its flattened last Conv2D, applies a LoRA-similar learned reduction, and concatenates.
  • v2 fixes v1's central flaw: both embeddings ignored class specifics. v2 embeds per class, keeps the 10 highest-activation classes, and adds the two reduced matrices instead of concatenating them.

Two unsupervised metrics

Distribution distance (DD) — the Euclidean distance between two datasets in (mean, |skew|, |kurtosis|) space over normalised pixel values:

$$\mathrm{DD}(i,j)=\sqrt{(\mu_i-\mu_j)^2+(s_i-s_j)^2+(\kappa_i-\kappa_j)^2}$$

Average class correlation (ACC) — mean pairwise correlation between per-class embeddings from an ImageNet-pretrained headless ResNet50 v2:

$$\mathrm{ACC}(i,j)=\frac{1}{nm}\sum_{k=1}^{n}\sum_{l=1}^{m}\mathrm{Cor}_{kl}$$

ACC separates similar from dissimilar transfer datasets far more sharply (0.501 vs. 0.276 — 45% lower for dissimilar) and explains fine-tune accuracy better than DD does.

Key findings

Pre-trained weights push dissimilar classes apart in latent space. Computing ACC from raw normalised pixels gives 0.507 for similar classes and 0.438 for dissimilar — a gap of just 0.06. Running the same computation through ImageNet-pretrained ResNet50 v2 weights widens that gap to 0.22 (0.49 vs. 0.27). The pre-trained weights are actively increasing class separation, which is a concrete answer to what transfer learning transfers.

A dataset's low-level pixel statistics explain much of the transfer effect. DD, computed from nothing but mean, skew, and kurtosis of pixel values, reaches an of 0.894.

Similar transfer datasets move weights less during fine-tuning. L2 distance between pre-trained and fine-tuned weights stays smaller across epochs for similar transfer datasets, corroborated by centered kernel alignment (CKA).

Honest limitations

  • The reported MSE figures come from a test subset that shares its fine-tuning dataset with the training subset. On a genuinely unseen fine-tuning dataset, v1 degrades badly — MSE 0.46, with accuracy predictions off by as much as 80%. Generalising across fine-tuning datasets needs either a v3 architecture or a TLDS spanning multiple fine-tuning datasets.
  • Test accuracy spans only ~15% from random initialisation to the best transfer-learned model, which is a narrow band in which to resolve differences.
  • Dissimilar transfer datasets scored about the same as random ones, so image diversity alone may matter less than expected here.
  • Everything is image classification on one fine-tuning task. Generalisation to other domains is untested.

Datasets

DatasetRole
Plant leaves (22 classes)Fine-tuning target
Plant diseases / plant_villageSimilar transfer source
Birds 450 speciesDissimilar transfer source
CIFAR-100, ImageNetRandom transfer sources

Download them all with the Kaggle API:

kaggle datasets download -p data/plant-diseases --unzip vipoooool/new-plant-diseases-dataset
kaggle datasets download -p data/plant-leaves --unzip csafrit2/plant-leaves-for-image-classification
kaggle datasets download -p data/bird-species --unzip gpiosenka/100-bird-species

Getting started

Developed with Python 3.10.

python -m venv venv
source venv/bin/activate
python -m pip install -r requirements.txt

The pipeline runs in three stages — build the TLDS, fine-tune, then train ChoiceNet on the result:

bash training/1.run_tl_training.sh # train TransferModel on each TL dataset
bash training/2.run_fine_tune.sh # fine-tune onto plant leaves, record accuracy
bash training/3.run_choicenet.sh # train TLDChoiceNet on the resulting TLDS

The shell scripts carry absolute paths from the original GPU machine (/data1/cs330/project/...). Point them at your own directories before running.

Monitor training:

tensorboard --logdir training # then open http://localhost:6006/

Code quality tooling

python -m pip install -r requirements-qa.txt
pre-commit install

Repository layout

data/ Dataset loading, preprocessing, and TLDS source configs
models/ TransferModel CNN and ChoiceNet v1/v2 architectures
training/ TLDS creation, fine-tuning, and ChoiceNet training entry points
embedding/ ResNet50 v2 dataset embedding and weight-matrix preprocessing
experiments/ Metric analysis — CKA, KS tests, correlation matrices, weight distance

Citation

@article{ning2026tldchoicenet,
title = {TLDChoiceNet: Quantitatively Choosing a Transfer Learning Dataset},
author = {Ning, Jing and Braza, James D.},
journal = {arXiv preprint},
year = {2026}
}

Acknowledgements

We thank Chelsea Finn and Daniel Zeng for helpful discussions and feedback, including ideas on ImageNet embedding and measuring weight distance across training.

About

Predict how well a transfer-learning dataset will work before fine-tuning on it. TLDChoiceNet cuts prediction MSE 5x, and an unsupervised class-correlation metric explains fine-tune accuracy with R2=0.97. Stanford CS 330.

Topics

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages

, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Strip utm_, fbclid, gclid, etc. from all links on page\n(function() {\n var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content',\n 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid',\n 'ref', 'ref_src', 'source', 'medium', 'campaign'];\n \n function cleanUrl(url) {\n try {\n var u = new URL(url, window.location.origin);\n var changed = false;\n trackingParams.forEach(function(p) {\n if (u.searchParams.has(p)) {\n u.searchParams.delete(p);\n changed = true;\n }\n });\n return changed ? u.toString() : url;\n } catch (e) {\n return url;\n }\n }\n \n function cleanLinks() {\n document.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n \n cleanLinks();\n \n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1) {\n if (node.tagName === 'A') cleanLinks();\n node.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Remove Tracking Parameters from Links"); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + '
Skip to content

Repository files navigation

TLDChoiceNet: Quantitatively Choosing a Transfer Learning Dataset

Predict how well a transfer-learning dataset will work — before spending the compute to fine-tune on it.

PosterPython 3.10TensorFlowCode style: black

Jing Ning, James D. Braza — Stanford University ·Paper (https://arxiv.org/abs/2608.09091)

Project poster

The original CS 330 poster — click for the full-resolution PDF. It covers TLDChoiceNet v1; the v2 architecture and the ACC metric below came later, so the figures in the poster differ from the results reported here.

Joint work. Originally developed at jamesbraza/cs330-project; this repository preserves the full commit history.


The problem

You have a small dataset for a target task. You have candidate pre-trained datasets A, B, and C — similar example counts, similar class counts — but compute to fine-tune on only one of them. Which do you pick?

Conventional wisdom says take the biggest, most diverse source. But no quantitative method existed to make that call. TLDChoiceNet predicts the post-fine-tuning test accuracy for each candidate, so the choice becomes a measurement instead of a guess.

Results

Model / metricResult
TLDChoiceNet v1 — test MSE0.154
TLDChoiceNet v2 — test MSE0.031 (5× lower)
Distribution distance (DD) vs. accuracy — 0.894
Average class correlation (ACC) vs. accuracy — 0.974

Two results worth separating:

The learned predictor. v2 cuts test MSE 5× over v1 by embedding both inputs per class rather than averaging across all of them, and by learning the post-embedding reduction rather than fixing it. It has ~1.15M trainable parameters against v1's ~314K.

The unsupervised metrics — no training required at all. Average class correlation reaches an of 0.974 against fine-tune accuracy. You can compute it from a pre-trained ResNet50 v2 forward pass and pick your transfer dataset without training anything.

Approach

The Transfer Learning Dataset (TLDS)

The core obstacle is that there is no dataset of "transfer learning dataset → resulting accuracy" pairs to learn from. So we built one. Each entry is a 3-tuple of (transfer learning dataset, transfer-learned model, resulting test accuracy), spanning four deliberately chosen regimes:

RegimeSourceSubsets
SimilarPlant diseases3 × 10 classes
DissimilarBird species3 × 10 classes
RandomCIFAR-100 and ImageNet40 × 10 classes
No transfer learningRandom initialisation (control)1

Each invocation of the generation script produces 57 data points; varying the seed yields more. Fine-tuning throughout is on a 22-class plant-leaves dataset.

TLDChoiceNet

Both versions take an embedded fine-tuning dataset and an embedded transfer-learning dataset/model, and regress to a single number — predicted test accuracy.

  • v1 embeds the fine-tuning dataset via PCA to 256 dimensions and the transfer-learned model via its flattened last Conv2D, applies a LoRA-similar learned reduction, and concatenates.
  • v2 fixes v1's central flaw: both embeddings ignored class specifics. v2 embeds per class, keeps the 10 highest-activation classes, and adds the two reduced matrices instead of concatenating them.

Two unsupervised metrics

Distribution distance (DD) — the Euclidean distance between two datasets in (mean, |skew|, |kurtosis|) space over normalised pixel values:

$$\mathrm{DD}(i,j)=\sqrt{(\mu_i-\mu_j)^2+(s_i-s_j)^2+(\kappa_i-\kappa_j)^2}$$

Average class correlation (ACC) — mean pairwise correlation between per-class embeddings from an ImageNet-pretrained headless ResNet50 v2:

$$\mathrm{ACC}(i,j)=\frac{1}{nm}\sum_{k=1}^{n}\sum_{l=1}^{m}\mathrm{Cor}_{kl}$$

ACC separates similar from dissimilar transfer datasets far more sharply (0.501 vs. 0.276 — 45% lower for dissimilar) and explains fine-tune accuracy better than DD does.

Key findings

Pre-trained weights push dissimilar classes apart in latent space. Computing ACC from raw normalised pixels gives 0.507 for similar classes and 0.438 for dissimilar — a gap of just 0.06. Running the same computation through ImageNet-pretrained ResNet50 v2 weights widens that gap to 0.22 (0.49 vs. 0.27). The pre-trained weights are actively increasing class separation, which is a concrete answer to what transfer learning transfers.

A dataset's low-level pixel statistics explain much of the transfer effect. DD, computed from nothing but mean, skew, and kurtosis of pixel values, reaches an of 0.894.

Similar transfer datasets move weights less during fine-tuning. L2 distance between pre-trained and fine-tuned weights stays smaller across epochs for similar transfer datasets, corroborated by centered kernel alignment (CKA).

Honest limitations

  • The reported MSE figures come from a test subset that shares its fine-tuning dataset with the training subset. On a genuinely unseen fine-tuning dataset, v1 degrades badly — MSE 0.46, with accuracy predictions off by as much as 80%. Generalising across fine-tuning datasets needs either a v3 architecture or a TLDS spanning multiple fine-tuning datasets.
  • Test accuracy spans only ~15% from random initialisation to the best transfer-learned model, which is a narrow band in which to resolve differences.
  • Dissimilar transfer datasets scored about the same as random ones, so image diversity alone may matter less than expected here.
  • Everything is image classification on one fine-tuning task. Generalisation to other domains is untested.

Datasets

DatasetRole
Plant leaves (22 classes)Fine-tuning target
Plant diseases / plant_villageSimilar transfer source
Birds 450 speciesDissimilar transfer source
CIFAR-100, ImageNetRandom transfer sources

Download them all with the Kaggle API:

kaggle datasets download -p data/plant-diseases --unzip vipoooool/new-plant-diseases-dataset
kaggle datasets download -p data/plant-leaves --unzip csafrit2/plant-leaves-for-image-classification
kaggle datasets download -p data/bird-species --unzip gpiosenka/100-bird-species

Getting started

Developed with Python 3.10.

python -m venv venv
source venv/bin/activate
python -m pip install -r requirements.txt

The pipeline runs in three stages — build the TLDS, fine-tune, then train ChoiceNet on the result:

bash training/1.run_tl_training.sh # train TransferModel on each TL dataset
bash training/2.run_fine_tune.sh # fine-tune onto plant leaves, record accuracy
bash training/3.run_choicenet.sh # train TLDChoiceNet on the resulting TLDS

The shell scripts carry absolute paths from the original GPU machine (/data1/cs330/project/...). Point them at your own directories before running.

Monitor training:

tensorboard --logdir training # then open http://localhost:6006/

Code quality tooling

python -m pip install -r requirements-qa.txt
pre-commit install

Repository layout

data/ Dataset loading, preprocessing, and TLDS source configs
models/ TransferModel CNN and ChoiceNet v1/v2 architectures
training/ TLDS creation, fine-tuning, and ChoiceNet training entry points
embedding/ ResNet50 v2 dataset embedding and weight-matrix preprocessing
experiments/ Metric analysis — CKA, KS tests, correlation matrices, weight distance

Citation

@article{ning2026tldchoicenet,
title = {TLDChoiceNet: Quantitatively Choosing a Transfer Learning Dataset},
author = {Ning, Jing and Braza, James D.},
journal = {arXiv preprint},
year = {2026}
}

Acknowledgements

We thank Chelsea Finn and Daniel Zeng for helpful discussions and feedback, including ideas on ImageNet embedding and measuring weight distance across training.

About

Predict how well a transfer-learning dataset will work before fine-tuning on it. TLDChoiceNet cuts prediction MSE 5x, and an unsupervised class-correlation metric explains fine-tune accuracy with R2=0.97. Stanford CS 330.

Topics

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages

, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Auto-enable theater mode on YouTube\n(function() {\n function tryTheater() {\n var btn = document.querySelector('button[aria-label=\"Theater mode\"], ytd-player #player button[title=\"Theater mode\"]');\n if (btn && !btn.classList.contains('activated')) {\n btn.click();\n }\n }\n \n // Try immediately\n tryTheater();\n \n // Try after navigation (SPA)\n var lastUrl = location.href;\n setInterval(function() {\n if (location.href !== lastUrl) {\n lastUrl = location.href;\n setTimeout(tryTheater, 500);\n }\n }, 1000);\n \n // Also try on player load\n var observer = new MutationObserver(tryTheater);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "YouTube Theater Mode Default"); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content

Repository files navigation

TLDChoiceNet: Quantitatively Choosing a Transfer Learning Dataset

Predict how well a transfer-learning dataset will work — before spending the compute to fine-tune on it.

PosterPython 3.10TensorFlowCode style: black

Jing Ning, James D. Braza — Stanford University ·Paper (https://arxiv.org/abs/2608.09091)

Project poster

The original CS 330 poster — click for the full-resolution PDF. It covers TLDChoiceNet v1; the v2 architecture and the ACC metric below came later, so the figures in the poster differ from the results reported here.

Joint work. Originally developed at jamesbraza/cs330-project; this repository preserves the full commit history.


The problem

You have a small dataset for a target task. You have candidate pre-trained datasets A, B, and C — similar example counts, similar class counts — but compute to fine-tune on only one of them. Which do you pick?

Conventional wisdom says take the biggest, most diverse source. But no quantitative method existed to make that call. TLDChoiceNet predicts the post-fine-tuning test accuracy for each candidate, so the choice becomes a measurement instead of a guess.

Results

Model / metricResult
TLDChoiceNet v1 — test MSE0.154
TLDChoiceNet v2 — test MSE0.031 (5× lower)
Distribution distance (DD) vs. accuracy — 0.894
Average class correlation (ACC) vs. accuracy — 0.974

Two results worth separating:

The learned predictor. v2 cuts test MSE 5× over v1 by embedding both inputs per class rather than averaging across all of them, and by learning the post-embedding reduction rather than fixing it. It has ~1.15M trainable parameters against v1's ~314K.

The unsupervised metrics — no training required at all. Average class correlation reaches an of 0.974 against fine-tune accuracy. You can compute it from a pre-trained ResNet50 v2 forward pass and pick your transfer dataset without training anything.

Approach

The Transfer Learning Dataset (TLDS)

The core obstacle is that there is no dataset of "transfer learning dataset → resulting accuracy" pairs to learn from. So we built one. Each entry is a 3-tuple of (transfer learning dataset, transfer-learned model, resulting test accuracy), spanning four deliberately chosen regimes:

RegimeSourceSubsets
SimilarPlant diseases3 × 10 classes
DissimilarBird species3 × 10 classes
RandomCIFAR-100 and ImageNet40 × 10 classes
No transfer learningRandom initialisation (control)1

Each invocation of the generation script produces 57 data points; varying the seed yields more. Fine-tuning throughout is on a 22-class plant-leaves dataset.

TLDChoiceNet

Both versions take an embedded fine-tuning dataset and an embedded transfer-learning dataset/model, and regress to a single number — predicted test accuracy.

  • v1 embeds the fine-tuning dataset via PCA to 256 dimensions and the transfer-learned model via its flattened last Conv2D, applies a LoRA-similar learned reduction, and concatenates.
  • v2 fixes v1's central flaw: both embeddings ignored class specifics. v2 embeds per class, keeps the 10 highest-activation classes, and adds the two reduced matrices instead of concatenating them.

Two unsupervised metrics

Distribution distance (DD) — the Euclidean distance between two datasets in (mean, |skew|, |kurtosis|) space over normalised pixel values:

$$\mathrm{DD}(i,j)=\sqrt{(\mu_i-\mu_j)^2+(s_i-s_j)^2+(\kappa_i-\kappa_j)^2}$$

Average class correlation (ACC) — mean pairwise correlation between per-class embeddings from an ImageNet-pretrained headless ResNet50 v2:

$$\mathrm{ACC}(i,j)=\frac{1}{nm}\sum_{k=1}^{n}\sum_{l=1}^{m}\mathrm{Cor}_{kl}$$

ACC separates similar from dissimilar transfer datasets far more sharply (0.501 vs. 0.276 — 45% lower for dissimilar) and explains fine-tune accuracy better than DD does.

Key findings

Pre-trained weights push dissimilar classes apart in latent space. Computing ACC from raw normalised pixels gives 0.507 for similar classes and 0.438 for dissimilar — a gap of just 0.06. Running the same computation through ImageNet-pretrained ResNet50 v2 weights widens that gap to 0.22 (0.49 vs. 0.27). The pre-trained weights are actively increasing class separation, which is a concrete answer to what transfer learning transfers.

A dataset's low-level pixel statistics explain much of the transfer effect. DD, computed from nothing but mean, skew, and kurtosis of pixel values, reaches an of 0.894.

Similar transfer datasets move weights less during fine-tuning. L2 distance between pre-trained and fine-tuned weights stays smaller across epochs for similar transfer datasets, corroborated by centered kernel alignment (CKA).

Honest limitations

  • The reported MSE figures come from a test subset that shares its fine-tuning dataset with the training subset. On a genuinely unseen fine-tuning dataset, v1 degrades badly — MSE 0.46, with accuracy predictions off by as much as 80%. Generalising across fine-tuning datasets needs either a v3 architecture or a TLDS spanning multiple fine-tuning datasets.
  • Test accuracy spans only ~15% from random initialisation to the best transfer-learned model, which is a narrow band in which to resolve differences.
  • Dissimilar transfer datasets scored about the same as random ones, so image diversity alone may matter less than expected here.
  • Everything is image classification on one fine-tuning task. Generalisation to other domains is untested.

Datasets

DatasetRole
Plant leaves (22 classes)Fine-tuning target
Plant diseases / plant_villageSimilar transfer source
Birds 450 speciesDissimilar transfer source
CIFAR-100, ImageNetRandom transfer sources

Download them all with the Kaggle API:

kaggle datasets download -p data/plant-diseases --unzip vipoooool/new-plant-diseases-dataset
kaggle datasets download -p data/plant-leaves --unzip csafrit2/plant-leaves-for-image-classification
kaggle datasets download -p data/bird-species --unzip gpiosenka/100-bird-species

Getting started

Developed with Python 3.10.

python -m venv venv
source venv/bin/activate
python -m pip install -r requirements.txt

The pipeline runs in three stages — build the TLDS, fine-tune, then train ChoiceNet on the result:

bash training/1.run_tl_training.sh # train TransferModel on each TL dataset
bash training/2.run_fine_tune.sh # fine-tune onto plant leaves, record accuracy
bash training/3.run_choicenet.sh # train TLDChoiceNet on the resulting TLDS

The shell scripts carry absolute paths from the original GPU machine (/data1/cs330/project/...). Point them at your own directories before running.

Monitor training:

tensorboard --logdir training # then open http://localhost:6006/

Code quality tooling

python -m pip install -r requirements-qa.txt
pre-commit install

Repository layout

data/ Dataset loading, preprocessing, and TLDS source configs
models/ TransferModel CNN and ChoiceNet v1/v2 architectures
training/ TLDS creation, fine-tuning, and ChoiceNet training entry points
embedding/ ResNet50 v2 dataset embedding and weight-matrix preprocessing
experiments/ Metric analysis — CKA, KS tests, correlation matrices, weight distance

Citation

@article{ning2026tldchoicenet,
title = {TLDChoiceNet: Quantitatively Choosing a Transfer Learning Dataset},
author = {Ning, Jing and Braza, James D.},
journal = {arXiv preprint},
year = {2026}
}

Acknowledgements

We thank Chelsea Finn and Daniel Zeng for helpful discussions and feedback, including ideas on ImageNet embedding and measuring weight distance across training.

About

Predict how well a transfer-learning dataset will work before fine-tuning on it. TLDChoiceNet cuts prediction MSE 5x, and an unsupervised class-correlation metric explains fine-tune accuracy with R2=0.97. Stanford CS 330.

Topics

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages

, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Remove or un-stick sticky/fixed headers that block content\n(function() {\n function unstick() {\n document.querySelectorAll('header, nav, [role=\"banner\"], .header, .navbar, .sticky, .fixed-top, [style*=\"position: fixed\"], [style*=\"position:sticky\"]').forEach(function(el) {\n if (el.style.position === 'fixed' || el.style.position === 'sticky' || \n getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') {\n el.style.position = 'static';\n el.style.top = 'auto';\n el.style.zIndex = 'auto';\n }\n });\n }\n \n unstick();\n \n var observer = new MutationObserver(unstick);\n observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] });\n})();", "Kill Sticky Headers"); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content

Repository files navigation

TLDChoiceNet: Quantitatively Choosing a Transfer Learning Dataset

Predict how well a transfer-learning dataset will work — before spending the compute to fine-tune on it.

PosterPython 3.10TensorFlowCode style: black

Jing Ning, James D. Braza — Stanford University ·Paper (https://arxiv.org/abs/2608.09091)

Project poster

The original CS 330 poster — click for the full-resolution PDF. It covers TLDChoiceNet v1; the v2 architecture and the ACC metric below came later, so the figures in the poster differ from the results reported here.

Joint work. Originally developed at jamesbraza/cs330-project; this repository preserves the full commit history.


The problem

You have a small dataset for a target task. You have candidate pre-trained datasets A, B, and C — similar example counts, similar class counts — but compute to fine-tune on only one of them. Which do you pick?

Conventional wisdom says take the biggest, most diverse source. But no quantitative method existed to make that call. TLDChoiceNet predicts the post-fine-tuning test accuracy for each candidate, so the choice becomes a measurement instead of a guess.

Results

Model / metricResult
TLDChoiceNet v1 — test MSE0.154
TLDChoiceNet v2 — test MSE0.031 (5× lower)
Distribution distance (DD) vs. accuracy — 0.894
Average class correlation (ACC) vs. accuracy — 0.974

Two results worth separating:

The learned predictor. v2 cuts test MSE 5× over v1 by embedding both inputs per class rather than averaging across all of them, and by learning the post-embedding reduction rather than fixing it. It has ~1.15M trainable parameters against v1's ~314K.

The unsupervised metrics — no training required at all. Average class correlation reaches an of 0.974 against fine-tune accuracy. You can compute it from a pre-trained ResNet50 v2 forward pass and pick your transfer dataset without training anything.

Approach

The Transfer Learning Dataset (TLDS)

The core obstacle is that there is no dataset of "transfer learning dataset → resulting accuracy" pairs to learn from. So we built one. Each entry is a 3-tuple of (transfer learning dataset, transfer-learned model, resulting test accuracy), spanning four deliberately chosen regimes:

RegimeSourceSubsets
SimilarPlant diseases3 × 10 classes
DissimilarBird species3 × 10 classes
RandomCIFAR-100 and ImageNet40 × 10 classes
No transfer learningRandom initialisation (control)1

Each invocation of the generation script produces 57 data points; varying the seed yields more. Fine-tuning throughout is on a 22-class plant-leaves dataset.

TLDChoiceNet

Both versions take an embedded fine-tuning dataset and an embedded transfer-learning dataset/model, and regress to a single number — predicted test accuracy.

  • v1 embeds the fine-tuning dataset via PCA to 256 dimensions and the transfer-learned model via its flattened last Conv2D, applies a LoRA-similar learned reduction, and concatenates.
  • v2 fixes v1's central flaw: both embeddings ignored class specifics. v2 embeds per class, keeps the 10 highest-activation classes, and adds the two reduced matrices instead of concatenating them.

Two unsupervised metrics

Distribution distance (DD) — the Euclidean distance between two datasets in (mean, |skew|, |kurtosis|) space over normalised pixel values:

$$\mathrm{DD}(i,j)=\sqrt{(\mu_i-\mu_j)^2+(s_i-s_j)^2+(\kappa_i-\kappa_j)^2}$$

Average class correlation (ACC) — mean pairwise correlation between per-class embeddings from an ImageNet-pretrained headless ResNet50 v2:

$$\mathrm{ACC}(i,j)=\frac{1}{nm}\sum_{k=1}^{n}\sum_{l=1}^{m}\mathrm{Cor}_{kl}$$

ACC separates similar from dissimilar transfer datasets far more sharply (0.501 vs. 0.276 — 45% lower for dissimilar) and explains fine-tune accuracy better than DD does.

Key findings

Pre-trained weights push dissimilar classes apart in latent space. Computing ACC from raw normalised pixels gives 0.507 for similar classes and 0.438 for dissimilar — a gap of just 0.06. Running the same computation through ImageNet-pretrained ResNet50 v2 weights widens that gap to 0.22 (0.49 vs. 0.27). The pre-trained weights are actively increasing class separation, which is a concrete answer to what transfer learning transfers.

A dataset's low-level pixel statistics explain much of the transfer effect. DD, computed from nothing but mean, skew, and kurtosis of pixel values, reaches an of 0.894.

Similar transfer datasets move weights less during fine-tuning. L2 distance between pre-trained and fine-tuned weights stays smaller across epochs for similar transfer datasets, corroborated by centered kernel alignment (CKA).

Honest limitations

  • The reported MSE figures come from a test subset that shares its fine-tuning dataset with the training subset. On a genuinely unseen fine-tuning dataset, v1 degrades badly — MSE 0.46, with accuracy predictions off by as much as 80%. Generalising across fine-tuning datasets needs either a v3 architecture or a TLDS spanning multiple fine-tuning datasets.
  • Test accuracy spans only ~15% from random initialisation to the best transfer-learned model, which is a narrow band in which to resolve differences.
  • Dissimilar transfer datasets scored about the same as random ones, so image diversity alone may matter less than expected here.
  • Everything is image classification on one fine-tuning task. Generalisation to other domains is untested.

Datasets

DatasetRole
Plant leaves (22 classes)Fine-tuning target
Plant diseases / plant_villageSimilar transfer source
Birds 450 speciesDissimilar transfer source
CIFAR-100, ImageNetRandom transfer sources

Download them all with the Kaggle API:

kaggle datasets download -p data/plant-diseases --unzip vipoooool/new-plant-diseases-dataset
kaggle datasets download -p data/plant-leaves --unzip csafrit2/plant-leaves-for-image-classification
kaggle datasets download -p data/bird-species --unzip gpiosenka/100-bird-species

Getting started

Developed with Python 3.10.

python -m venv venv
source venv/bin/activate
python -m pip install -r requirements.txt

The pipeline runs in three stages — build the TLDS, fine-tune, then train ChoiceNet on the result:

bash training/1.run_tl_training.sh # train TransferModel on each TL dataset
bash training/2.run_fine_tune.sh # fine-tune onto plant leaves, record accuracy
bash training/3.run_choicenet.sh # train TLDChoiceNet on the resulting TLDS

The shell scripts carry absolute paths from the original GPU machine (/data1/cs330/project/...). Point them at your own directories before running.

Monitor training:

tensorboard --logdir training # then open http://localhost:6006/

Code quality tooling

python -m pip install -r requirements-qa.txt
pre-commit install

Repository layout

data/ Dataset loading, preprocessing, and TLDS source configs
models/ TransferModel CNN and ChoiceNet v1/v2 architectures
training/ TLDS creation, fine-tuning, and ChoiceNet training entry points
embedding/ ResNet50 v2 dataset embedding and weight-matrix preprocessing
experiments/ Metric analysis — CKA, KS tests, correlation matrices, weight distance

Citation

@article{ning2026tldchoicenet,
title = {TLDChoiceNet: Quantitatively Choosing a Transfer Learning Dataset},
author = {Ning, Jing and Braza, James D.},
journal = {arXiv preprint},
year = {2026}
}

Acknowledgements

We thank Chelsea Finn and Daniel Zeng for helpful discussions and feedback, including ideas on ImageNet embedding and measuring weight distance across training.

About

Predict how well a transfer-learning dataset will work before fine-tuning on it. TLDChoiceNet cuts prediction MSE 5x, and an unsupervised class-correlation metric explains fine-tune accuracy with R2=0.97. Stanford CS 330.

Topics

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages

, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Universal Dark Mode - works on any site\n(function() {\n var enabled = true;\n \n function applyDarkMode() {\n if (!enabled) return;\n \n // Create style element if it doesn't exist\n var style = document.getElementById('universal-dark-mode-style');\n if (!style) {\n style = document.createElement('style');\n style.id = 'universal-dark-mode-style';\n document.head.appendChild(style);\n }\n \n // Dark mode CSS - inverts colors but preserves images/video\n style.textContent = '\n /* Invert everything except media */\n html {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #1a1a2e !important;\n }\n \n /* Restore images, videos, iframes, canvas */\n img, video, iframe, canvas, svg, picture, [style*=\"background-image\"] {\n filter: invert(1) hue-rotate(180deg) !important;\n }\n \n /* Preserve specific elements that should not be inverted */\n .no-dark-mode, .no-dark-mode *,\n [data-theme=\"light\"], [data-theme=\"light\"],\n .ace_editor, .ace_editor *,\n .CodeMirror, .CodeMirror *,\n .monaco-editor, .monaco-editor *,\n .markdown-body pre, .markdown-body pre *,\n .highlight, .highlight *,\n pre code, pre code * {\n filter: none !important;\n }\n \n /* Fix common UI elements */\n .modal, .popup, .dropdown-menu, .tooltip, .popover {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #2d2d44 !important;\n border-color: #444 !important;\n }\n \n /* Scrollbars */\n ::-webkit-scrollbar { background: #1a1a2e !important; }\n ::-webkit-scrollbar-thumb { background: #444 !important; }\n ::-webkit-scrollbar-thumb:hover { background: #555 !important; }\n \n /* Selection */\n ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ';\n }\n \n function removeDarkMode() {\n var style = document.getElementById('universal-dark-mode-style');\n if (style) style.remove();\n }\n \n // Toggle with Alt+Shift+D\n document.addEventListener('keydown', function(e) {\n if (e.altKey && e.shiftKey && e.key === 'D') {\n e.preventDefault();\n enabled = !enabled;\n if (enabled) {\n applyDarkMode();\n console.log('[Universal Dark Mode] Enabled');\n } else {\n removeDarkMode();\n console.log('[Universal Dark Mode] Disabled');\n }\n }\n });\n \n // Apply on load\n applyDarkMode();\n \n // Re-apply on dynamic content\n var observer = new MutationObserver(function(mutations) {\n if (enabled && !document.getElementById('universal-dark-mode-style')) {\n applyDarkMode();\n }\n });\n observer.observe(document.head, { childList: true });\n \n console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle');\n})();", "Universal Dark Mode"); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })();
Skip to content

Repository files navigation

TLDChoiceNet: Quantitatively Choosing a Transfer Learning Dataset

Predict how well a transfer-learning dataset will work — before spending the compute to fine-tune on it.

PosterPython 3.10TensorFlowCode style: black

Jing Ning, James D. Braza — Stanford University ·Paper (https://arxiv.org/abs/2608.09091)

Project poster

The original CS 330 poster — click for the full-resolution PDF. It covers TLDChoiceNet v1; the v2 architecture and the ACC metric below came later, so the figures in the poster differ from the results reported here.

Joint work. Originally developed at jamesbraza/cs330-project; this repository preserves the full commit history.


The problem

You have a small dataset for a target task. You have candidate pre-trained datasets A, B, and C — similar example counts, similar class counts — but compute to fine-tune on only one of them. Which do you pick?

Conventional wisdom says take the biggest, most diverse source. But no quantitative method existed to make that call. TLDChoiceNet predicts the post-fine-tuning test accuracy for each candidate, so the choice becomes a measurement instead of a guess.

Results

Model / metricResult
TLDChoiceNet v1 — test MSE0.154
TLDChoiceNet v2 — test MSE0.031 (5× lower)
Distribution distance (DD) vs. accuracy — 0.894
Average class correlation (ACC) vs. accuracy — 0.974

Two results worth separating:

The learned predictor. v2 cuts test MSE 5× over v1 by embedding both inputs per class rather than averaging across all of them, and by learning the post-embedding reduction rather than fixing it. It has ~1.15M trainable parameters against v1's ~314K.

The unsupervised metrics — no training required at all. Average class correlation reaches an of 0.974 against fine-tune accuracy. You can compute it from a pre-trained ResNet50 v2 forward pass and pick your transfer dataset without training anything.

Approach

The Transfer Learning Dataset (TLDS)

The core obstacle is that there is no dataset of "transfer learning dataset → resulting accuracy" pairs to learn from. So we built one. Each entry is a 3-tuple of (transfer learning dataset, transfer-learned model, resulting test accuracy), spanning four deliberately chosen regimes:

RegimeSourceSubsets
SimilarPlant diseases3 × 10 classes
DissimilarBird species3 × 10 classes
RandomCIFAR-100 and ImageNet40 × 10 classes
No transfer learningRandom initialisation (control)1

Each invocation of the generation script produces 57 data points; varying the seed yields more. Fine-tuning throughout is on a 22-class plant-leaves dataset.

TLDChoiceNet

Both versions take an embedded fine-tuning dataset and an embedded transfer-learning dataset/model, and regress to a single number — predicted test accuracy.

  • v1 embeds the fine-tuning dataset via PCA to 256 dimensions and the transfer-learned model via its flattened last Conv2D, applies a LoRA-similar learned reduction, and concatenates.
  • v2 fixes v1's central flaw: both embeddings ignored class specifics. v2 embeds per class, keeps the 10 highest-activation classes, and adds the two reduced matrices instead of concatenating them.

Two unsupervised metrics

Distribution distance (DD) — the Euclidean distance between two datasets in (mean, |skew|, |kurtosis|) space over normalised pixel values:

$$\mathrm{DD}(i,j)=\sqrt{(\mu_i-\mu_j)^2+(s_i-s_j)^2+(\kappa_i-\kappa_j)^2}$$

Average class correlation (ACC) — mean pairwise correlation between per-class embeddings from an ImageNet-pretrained headless ResNet50 v2:

$$\mathrm{ACC}(i,j)=\frac{1}{nm}\sum_{k=1}^{n}\sum_{l=1}^{m}\mathrm{Cor}_{kl}$$

ACC separates similar from dissimilar transfer datasets far more sharply (0.501 vs. 0.276 — 45% lower for dissimilar) and explains fine-tune accuracy better than DD does.

Key findings

Pre-trained weights push dissimilar classes apart in latent space. Computing ACC from raw normalised pixels gives 0.507 for similar classes and 0.438 for dissimilar — a gap of just 0.06. Running the same computation through ImageNet-pretrained ResNet50 v2 weights widens that gap to 0.22 (0.49 vs. 0.27). The pre-trained weights are actively increasing class separation, which is a concrete answer to what transfer learning transfers.

A dataset's low-level pixel statistics explain much of the transfer effect. DD, computed from nothing but mean, skew, and kurtosis of pixel values, reaches an of 0.894.

Similar transfer datasets move weights less during fine-tuning. L2 distance between pre-trained and fine-tuned weights stays smaller across epochs for similar transfer datasets, corroborated by centered kernel alignment (CKA).

Honest limitations

  • The reported MSE figures come from a test subset that shares its fine-tuning dataset with the training subset. On a genuinely unseen fine-tuning dataset, v1 degrades badly — MSE 0.46, with accuracy predictions off by as much as 80%. Generalising across fine-tuning datasets needs either a v3 architecture or a TLDS spanning multiple fine-tuning datasets.
  • Test accuracy spans only ~15% from random initialisation to the best transfer-learned model, which is a narrow band in which to resolve differences.
  • Dissimilar transfer datasets scored about the same as random ones, so image diversity alone may matter less than expected here.
  • Everything is image classification on one fine-tuning task. Generalisation to other domains is untested.

Datasets

DatasetRole
Plant leaves (22 classes)Fine-tuning target
Plant diseases / plant_villageSimilar transfer source
Birds 450 speciesDissimilar transfer source
CIFAR-100, ImageNetRandom transfer sources

Download them all with the Kaggle API:

kaggle datasets download -p data/plant-diseases --unzip vipoooool/new-plant-diseases-dataset
kaggle datasets download -p data/plant-leaves --unzip csafrit2/plant-leaves-for-image-classification
kaggle datasets download -p data/bird-species --unzip gpiosenka/100-bird-species

Getting started

Developed with Python 3.10.

python -m venv venv
source venv/bin/activate
python -m pip install -r requirements.txt

The pipeline runs in three stages — build the TLDS, fine-tune, then train ChoiceNet on the result:

bash training/1.run_tl_training.sh # train TransferModel on each TL dataset
bash training/2.run_fine_tune.sh # fine-tune onto plant leaves, record accuracy
bash training/3.run_choicenet.sh # train TLDChoiceNet on the resulting TLDS

The shell scripts carry absolute paths from the original GPU machine (/data1/cs330/project/...). Point them at your own directories before running.

Monitor training:

tensorboard --logdir training # then open http://localhost:6006/

Code quality tooling

python -m pip install -r requirements-qa.txt
pre-commit install

Repository layout

data/ Dataset loading, preprocessing, and TLDS source configs
models/ TransferModel CNN and ChoiceNet v1/v2 architectures
training/ TLDS creation, fine-tuning, and ChoiceNet training entry points
embedding/ ResNet50 v2 dataset embedding and weight-matrix preprocessing
experiments/ Metric analysis — CKA, KS tests, correlation matrices, weight distance

Citation

@article{ning2026tldchoicenet,
title = {TLDChoiceNet: Quantitatively Choosing a Transfer Learning Dataset},
author = {Ning, Jing and Braza, James D.},
journal = {arXiv preprint},
year = {2026}
}

Acknowledgements

We thank Chelsea Finn and Daniel Zeng for helpful discussions and feedback, including ideas on ImageNet embedding and measuring weight distance across training.

About

Predict how well a transfer-learning dataset will work before fine-tuning on it. TLDChoiceNet cuts prediction MSE 5x, and an unsupervised class-correlation metric explains fine-tune accuracy with R2=0.97. Stanford CS 330.

Topics

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages