TP4: Feature Learning in Infinite-Width Neural Networks
When to Use
- Implementing infinite-width neural network experiments (GP, NTK, muP/feature learning limits)
- Replicating Word2Vec experiments with infinite-width models
- Running MAML (Model-Agnostic Meta-Learning) with infinite-width networks on Omniglot
- Studying the Tensor Programs series of papers empirically
- Comparing finite vs. infinite-width neural network behavior
- Keywords: infinite-width, feature learning, NTK, GP, muP, MAML, Word2Vec, Tensor Programs, meta-learning
Quick Reference
Installation / Setup
Prerequisites
- Python 3.x
- C compiler (for Word2Vec C source)
- PyTorch
MAML Experiment Setup
cd TP4MAML
pip install -r requirements.txt
cd meta
pip install -r requirements.txt
Word2Vec Experiment Setup
cd Word2Vec
# Build C binaries
make
# Download and prepare text8 dataset
bash scripts/create-text8-data.sh
# Download and prepare fil9 dataset
bash scripts/create-fil9-data.sh
Core Features
- InfGP1LP: Infinite-width Gaussian Process limit for a 1-hidden-layer perceptron
- FinGP1LP: Finite-width GP baseline for comparison
- InfNTK1LP: Infinite-width NTK limit for a 1-hidden-layer perceptron
- InfSGD: Custom SGD optimizer for infinite-width networks with proper scaling
- InfMultiStepLR: Learning rate scheduler compatible with InfSGD
- InfMAML: Infinite-width MAML metalearner for Omniglot few-shot classification
- CachedOmniglot: Efficient cached Omniglot dataset loader
- Word2Vec C implementation: Modified word2vec supporting infinite-width training modes
Usage Examples
Running All MAML Experiments
cd TP4MAML
bash train_all.sh
Training MAML (finite width)
cd TP4MAML/meta
python train.py --dataset omniglot --num-ways 5 --num-shots 1
Training Infinite-Width MAML
cd TP4MAML/meta
python train.py --dataset omniglot --num-ways 5 --num-shots 1 --inf
Word2Vec Training (text8, standard)
cd Word2Vec
bash scripts/train-text8.sh
Word2Vec Training (text8, infinite-width)
cd Word2Vec
bash scripts/train-text8-inf.sh
Word2Vec Evaluation
cd Word2Vec
bash scripts/evaluate.sh
Key APIs / Models
TP4MAML/inf/inf1lp.py
InfGP1LP — Infinite GP limit 1-layer perceptron
FinGP1LP — Finite GP baseline
InfNTK1LP — Infinite NTK limit 1-layer perceptron
TP4MAML/inf/optim.py
InfSGD(params, lr, ...) — SGD optimizer scaled for infinite-width networks
InfMultiStepLR(optimizer, milestones, gamma) — LR scheduler for InfSGD
TP4MAML/inf/utils.py
safe_sqrt(arr, eps) — Numerically stable square root
safe_acos(arr, eps) — Numerically stable arccos
F00ReLU(c, v, v2) — ReLU kernel function used in GP/NTK computations
MyLinear — Custom linear layer with infinite-width scaling
TP4MAML/meta/maml/metalearners/infmaml.py
InfMAML — Meta-learner implementing MAML for infinite-width networks
TP4MAML/meta/maml/metalearners/maml.py
MAML — Standard MAML meta-learner
TP4MAML/meta/cached_omniglot.py
CachedOmniglot — Omniglot dataset with caching
OmniglotClassDataset — Per-class Omniglot dataset
Omniglot — Base Omniglot loader
TP4MAML/inf/dynamicarray.py
DynArr — Dynamic array for storing activations during infinite-width forward passes
CycArr — Cyclic array variant
Common Patterns & Best Practices
- Use
--inf flag in train.py to switch between finite and infinite-width MAML
- The infinite-width models do not store explicit weights; instead they accumulate kernel computations
- For Word2Vec, the
train-*-inf.sh scripts set hyperparameters appropriate for infinite-width training
- Always build the C binaries before running Word2Vec experiments (
make in Word2Vec/)
- The
train_all.sh script in TP4MAML/ runs all configurations sequentially for full replication
Demo Scripts
scripts/inf_network_demo.py
#!/usr/bin/env python3
"""
TP4 Feature Learning - Infinite-Width Network Demo
Demonstrates usage of the TP4MAML inf module:
- InfGP1LP, InfNTK1LP for infinite-width 1-hidden-layer perceptrons
- InfSGD optimizer
- Utility functions (safe_sqrt, safe_acos, F00ReLU)
Requires: torch, numpy
Run from the repo root: python scripts/inf_network_demo.py
"""
import sys
import os
# Add TP4MAML to path so we can import the inf module
sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..', 'TP4MAML'))
import torch
import torch.nn as nn
import numpy as np
def demo_utils():
"""Demonstrate utility functions from TP4MAML/inf/utils.py"""
try:
from inf.utils import safe_sqrt, safe_acos, F00ReLU, MyLinear
print("=== Utility Functions ===")
# safe_sqrt: numerically stable square root
arr = torch.tensor([-1e-10, 0.0, 1.0, 4.0])
result = safe_sqrt(arr, eps=1e-6)
print(f"safe_sqrt({arr.tolist()}) = {result.tolist()}")
# safe_acos: numerically stable arccos
arr2 = torch.tensor([-1.0 - 1e-10, -1.0, 0.0, 1.0, 1.0 + 1e-10])
result2 = safe_acos(arr2, eps=1e-6)
print(f"safe_acos(clipped) = {result2.tolist()}")
# F00ReLU: ReLU arc-cosine kernel
# c: cosine similarity, v: variance1, v2: variance2
c = torch.tensor([0.5])
v = torch.tensor([1.0])
v2 = torch.tensor([1.0])
k = F00ReLU(c, v, v2)
print(f"F00ReLU(c=0.5, v=1, v2=1) = {k.item():.4f}")
# MyLinear: custom linear layer with infinite-width scaling
linear = MyLinear(in_features=10, out_features=5)
x = torch.randn(3, 10)
out = linear(x)
print(f"MyLinear(10->5) output shape: {out.shape}")
except ImportError as e:
print(f"Could not import inf.utils (run from repo root with TP4MAML in path): {e}")
def demo_dynamic_arrays():
"""Demonstrate DynArr and CycArr from TP4MAML/inf/dynamicarray.py"""
try:
from inf.dynamicarray import DynArr, CycArr
print("\n=== Dynamic Arrays ===")
# DynArr: growing array for storing activations
darr = DynArr()
for i in range(5):
darr.append(torch.randn(3, 4))
print(f"DynArr length after 5 appends: {len(darr)}")
# CycArr: cyclic (ring buffer) array
carr = CycArr(capacity=3)
for i in range(6):
carr.append(torch.tensor([float(i)]))
print(f"CycArr (capacity=3) last 3 values: {[carr[i].item() for i in range(len(carr))]}")
except ImportError as e:
print(f"Could not import inf.dynamicarray: {e}")
def demo_inf_models():
"""Demonstrate InfGP1LP, FinGP1LP, InfNTK1LP from TP4MAML/inf/inf1lp.py"""
try:
from inf.inf1lp import InfGP1LP, FinGP1LP, InfNTK1LP
print("\n=== Infinite-Width 1-Layer Perceptron Models ===")
input_dim = 16
output_dim = 5
batch_size = 8
x_train = torch.randn(batch_size, input_dim)
y_train = torch.randint(0, output_dim, (batch_size,))
x_test = torch.randn(4, input_dim)
# InfGP1LP: Infinite-width Gaussian Process limit
print("Building InfGP1LP...")
gp_model = InfGP1LP(input_dim=input_dim, output_dim=output_dim)
print(f" InfGP1LP created: {type(gp_model).__name__}")
# InfNTK1LP: Infinite-width NTK limit
print("Building InfNTK1LP...")
ntk_model = InfNTK1LP(input_dim=input_dim, output_dim=output_dim)
print(f" InfNTK1LP created: {type(ntk_model).__name__}")
# FinGP1LP: Finite baseline
print("Building FinGP1LP...")
fin_model = FinGP1LP(input_dim=input_dim, output_dim=output_dim, width=256)
print(f" FinGP1LP created: {type(fin_model).__name__}")
# Forward pass on finite model
logits = fin_model(x_test)
print(f" FinGP1LP forward output shape: {logits.shape}")
except ImportError as e:
print(f"Could not import inf.inf1lp: {e}")
except Exception as e:
print(f"Error in inf model demo: {e}")
def demo_inf_sgd():
"""Demonstrate InfSGD optimizer from TP4MAML/inf/optim.py"""
try:
from inf.optim import InfSGD, InfMultiStepLR
print("\n=== InfSGD Optimizer ===")
# Simple model to optimize
model = nn.Linear(10, 5)
# InfSGD: SGD with infinite-width scaling
optimizer = InfSGD(model.parameters(), lr=0.01, momentum=0.9)
print(f"InfSGD created with lr=0.01, momentum=0.9")
# InfMultiStepLR scheduler
scheduler = InfMultiStepLR(optimizer, milestones=[10, 20], gamma=0.1)
print(f"InfMultiStepLR created with milestones=[10, 20], gamma=0.1")
# Simulate a few training steps
x = torch.randn(4, 10)
y = torch.randint(0, 5, (4,))
criterion = nn.CrossEntropyLoss()
for step in range(3):
optimizer.zero_grad()
out = model(x)
loss = criterion(out, y)
loss.backward()
optimizer.step()
scheduler.step()
print(f" Step {step+1}: loss={loss.item():.4f}, lr={scheduler.get_last_lr()}")
except ImportError as e:
print(f"Could not import inf.optim: {e}")
except Exception as e:
print(f"Error in InfSGD demo: {e}")
def demo_maml_structure():
"""Show the MAML metalearner interface (TP4MAML/meta/maml/metalearners/)"""
try:
sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..', 'TP4MAML', 'meta'))
from maml.metalearners.maml import MAML
print("\n=== MAML Metalearner ===")
print(f"MAML class imported: {MAML}")
print(" MAML implements standard Model-Agnostic Meta-Learning.")
print(" Use train.py --inf flag to switch to InfMAML.")
except ImportError as e:
print(f"\nCould not import MAML metalearner: {e}")
if __name__ == "__main__":
print("TP4 Feature Learning in Infinite-Width Neural Networks - Demo\n")
demo_utils()
demo_dynamic_arrays()
demo_inf_models()
demo_inf_sgd()
demo_maml_structure()
print("\nDemo complete.")
scripts/maml_training_demo.py
#!/usr/bin/env python3
"""
TP4 MAML Training Demo
Shows how to invoke MAML training programmatically (mirroring train.py usage).
Demonstrates the dataset loading and metalearner API.
Requires: torch, torchmeta (see TP4MAML/meta/requirements.txt)
Run from TP4MAML/meta/: python ../../scripts/maml_training_demo.py
"""
import sys
import os
import argparse
# Adjust paths for running from repo root
MAML_META_PATH = os.path.join(os.path.dirname(__file__), '..', 'TP4MAML', 'meta')
MAML_INF_PATH = os.path.join(os.path.dirname(__file__), '..', 'TP4MAML')
sys.path.insert(0, MAML_META_PATH)
sys.path.insert(0, MAML_INF_PATH)
def build_omniglot_dataset(data_folder: str, num_ways: int = 5, num_shots: int = 1,
num_shots_test: int = 15):
"""
Build Omniglot meta-learning dataset using CachedOmniglot.
Args:
data_folder: Path to Omniglot data directory.
num_ways: Number of classes per episode (N-way).
num_shots: Number of support examples per class (K-shot).
num_shots_test: Number of query examples per class.
Returns:
Tuple of (meta_train_dataset, meta_val_dataset, meta_test_dataset)
"""
try:
from cached_omniglot import CachedOmniglot
import torchmeta
from torchmeta.transforms import ClassSplitter, Categorical
from torchvision.transforms import Compose, Resize, ToTensor
transform = Compose([Resize(28), ToTensor()])
meta_train = CachedOmniglot(
data_folder,
num_classes_per_task=num_ways,
transform=transform,
target_transform=Categorical(num_ways),
class_augmentations=[torchmeta.transforms.Rotation([90, 180, 270])],
meta_train=True,
dataset_transform=ClassSplitter(
shuffle=True,
num_support_per_class=num_shots,
num_query_per_class=num_shots_test
)
)
print(f"CachedOmniglot meta-train: {len(meta_train)} tasks")
return meta_train
except ImportError as e:
print(f"Could not build Omniglot dataset (missing torchmeta or data): {e}")
return None
except Exception as e:
print(f"Dataset error: {e}")
return None
def demonstrate_metalearner_api():
"""
Show the API surface of MAML and InfMAML metalearners.
"""
print("=== MAML / InfMAML API ===\n")
try:
from maml.metalearners.maml import MAML
print(f"MAML class: {MAML.__module__}.{MAML.__name__}")
print(f" __init__ signature: see references/api_reference.md")
print(f" Key methods: train(), evaluate(), get_outer_loss()")
except ImportError as e:
print(f"MAML import failed: {e}")
try:
from maml.metalearners.infmaml import InfMAML
print(f"\nInfMAML class: {InfMAML.__module__}.{InfMAML.__name__}")
print(f" Extends MAML for infinite-width networks using InfGP1LP/InfNTK1LP")
except ImportError as e:
print(f"InfMAML import failed: {e}")
try:
from maml.metalearners.meta_sgd import MetaSGD
print(f"\nMetaSGD class: {MetaSGD.__module__}.{MetaSGD.__name__}")
print(f" Meta-SGD variant with per-parameter learned learning rates")
except ImportError as e:
print(f"MetaSGD import failed: {e}")
def show_training_command_equivalents():
"""
Print the equivalent train.py CLI commands for common configurations.
"""
print("\n=== Equivalent train.py Commands ===\n")
configs = [
{
"description": "5-way 1-shot Omniglot, finite MAML",
"cmd": "python train.py --dataset omniglot --num-ways 5 --num-shots 1 --num-steps 5"
},
{
"description": "5-way 1-shot Omniglot, infinite-width MAML (GP limit)",
"cmd": "python train.py --dataset omniglot --num-ways 5 --num-shots 1 --num-steps 5 --inf --use-gp"
},
{
"description": "5-way 5-shot Omniglot, infinite-width MAML (NTK limit)",
"cmd": "python train.py --dataset omniglot --num-ways 5 --num-shots 5 --num-steps 5 --inf"
},
{
"description": "Run all experiments (uses train_all.sh)",
"cmd": "cd TP4MAML && bash train_all.sh"
},
]
for cfg in configs:
print(f"# {cfg['description']}")
print(f" {cfg['cmd']}\n")
def show_word2vec_commands():
"""
Show Word2Vec experiment commands.
"""
print("=== Word2Vec Experiment Commands ===\n")
steps = [
("Build C binaries", "cd Word2Vec && make"),
("Prepare text8 data", "bash Word2Vec/scripts/create-text8-data.sh"),
("Prepare fil9 data", "bash Word2Vec/scripts/create-fil9-data.sh"),
("Train text8 (finite)", "bash Word2Vec/scripts/train-text8.sh"),
("Train text8 (infinite)", "bash Word2Vec/scripts/train-text8-inf.sh"),
("Train fil9 (finite)", "bash Word2Vec/scripts/train-fil9.sh"),
("Train fil9 (infinite)", "bash Word2Vec/scripts/train-fil9-inf.sh"),
("Evaluate embeddings", "bash Word2Vec/scripts/evaluate.sh"),
]
for name, cmd in steps:
print(f"# {name}")
print(f" {cmd}\n")
if __name__ == "__main__":
print("TP4 MAML Training Demo\n")
demonstrate_metalearner_api()
show_training_command_equivalents()
show_word2vec_commands()
# Optionally build dataset if data path provided
if len(sys.argv) > 1:
data_folder = sys.argv[1]
print(f"\n=== Building Omniglot Dataset from {data_folder} ===")
dataset = build_omniglot_dataset(data_folder, num_ways=5, num_shots=1)
if dataset is not None:
print("Dataset built successfully.")
else:
print("\nTip: Pass a data folder path as argument to test dataset loading.")
print(" python maml_training_demo.py /path/to/omniglot/data")
1---2name: tp4-feature-learning3description: Use this skill when working with infinite-width neural networks for feature learning, replicating Word2Vec or MAML experiments from the Tensor Programs series (TP4), or implementing infinite-width limits (GP, NTK, muP) for meta-learning and word embedding tasks.4---56# TP4: Feature Learning in Infinite-Width Neural Networks78## When to Use9- Implementing infinite-width neural network experiments (GP, NTK, muP/feature learning limits)10- Replicating Word2Vec experiments with infinite-width models11- Running MAML (Model-Agnostic Meta-Learning) with infinite-width networks on Omniglot12- Studying the Tensor Programs series of papers empirically13- Comparing finite vs. infinite-width neural network behavior14- Keywords: infinite-width, feature learning, NTK, GP, muP, MAML, Word2Vec, Tensor Programs, meta-learning1516## Quick Reference17- **Paper:** https://arxiv.org/abs/2011.1452218- **Repo:** https://github.com/edwardjhu/TP419- **GP limit code:** https://github.com/thegregyang/GP4A20- **NTK limit code:** https://github.com/thegregyang/NTK4A21- **Prior TP papers:** [TP0](http://arxiv.org/abs/1902.04760), [TP1](http://arxiv.org/abs/1910.12478), [TP2](http://arxiv.org/abs/2006.14548), [TP3](http://arxiv.org/abs/2009.10685)2223## Installation / Setup2425### Prerequisites26- Python 3.x27- C compiler (for Word2Vec C source)28- PyTorch2930### MAML Experiment Setup31```bash32cd TP4MAML33pip install -r requirements.txt34cd meta35pip install -r requirements.txt36```3738### Word2Vec Experiment Setup39```bash40cd Word2Vec41# Build C binaries42make43# Download and prepare text8 dataset44bash scripts/create-text8-data.sh45# Download and prepare fil9 dataset46bash scripts/create-fil9-data.sh47```4849## Core Features5051- **InfGP1LP:** Infinite-width Gaussian Process limit for a 1-hidden-layer perceptron52- **FinGP1LP:** Finite-width GP baseline for comparison53- **InfNTK1LP:** Infinite-width NTK limit for a 1-hidden-layer perceptron54- **InfSGD:** Custom SGD optimizer for infinite-width networks with proper scaling55- **InfMultiStepLR:** Learning rate scheduler compatible with InfSGD56- **InfMAML:** Infinite-width MAML metalearner for Omniglot few-shot classification57- **CachedOmniglot:** Efficient cached Omniglot dataset loader58- **Word2Vec C implementation:** Modified word2vec supporting infinite-width training modes5960## Usage Examples6162### Running All MAML Experiments63```bash64cd TP4MAML65bash train_all.sh66```6768### Training MAML (finite width)69```bash70cd TP4MAML/meta71python train.py --dataset omniglot --num-ways 5 --num-shots 172```7374### Training Infinite-Width MAML75```bash76cd TP4MAML/meta77python train.py --dataset omniglot --num-ways 5 --num-shots 1 --inf78```7980### Word2Vec Training (text8, standard)81```bash82cd Word2Vec83bash scripts/train-text8.sh84```8586### Word2Vec Training (text8, infinite-width)87```bash88cd Word2Vec89bash scripts/train-text8-inf.sh90```9192### Word2Vec Evaluation93```bash94cd Word2Vec95bash scripts/evaluate.sh96```9798## Key APIs / Models99100### `TP4MAML/inf/inf1lp.py`101- `InfGP1LP` — Infinite GP limit 1-layer perceptron102- `FinGP1LP` — Finite GP baseline103- `InfNTK1LP` — Infinite NTK limit 1-layer perceptron104105### `TP4MAML/inf/optim.py`106- `InfSGD(params, lr, ...)` — SGD optimizer scaled for infinite-width networks107- `InfMultiStepLR(optimizer, milestones, gamma)` — LR scheduler for InfSGD108109### `TP4MAML/inf/utils.py`110- `safe_sqrt(arr, eps)` — Numerically stable square root111- `safe_acos(arr, eps)` — Numerically stable arccos112- `F00ReLU(c, v, v2)` — ReLU kernel function used in GP/NTK computations113- `MyLinear` — Custom linear layer with infinite-width scaling114115### `TP4MAML/meta/maml/metalearners/infmaml.py`116- `InfMAML` — Meta-learner implementing MAML for infinite-width networks117118### `TP4MAML/meta/maml/metalearners/maml.py`119- `MAML` — Standard MAML meta-learner120121### `TP4MAML/meta/cached_omniglot.py`122- `CachedOmniglot` — Omniglot dataset with caching123- `OmniglotClassDataset` — Per-class Omniglot dataset124- `Omniglot` — Base Omniglot loader125126### `TP4MAML/inf/dynamicarray.py`127- `DynArr` — Dynamic array for storing activations during infinite-width forward passes128- `CycArr` — Cyclic array variant129130## Common Patterns & Best Practices131132- Use `--inf` flag in `train.py` to switch between finite and infinite-width MAML133- The infinite-width models do not store explicit weights; instead they accumulate kernel computations134- For Word2Vec, the `train-*-inf.sh` scripts set hyperparameters appropriate for infinite-width training135- Always build the C binaries before running Word2Vec experiments (`make` in `Word2Vec/`)136- The `train_all.sh` script in `TP4MAML/` runs all configurations sequentially for full replication137138## Demo Scripts139140### `scripts/inf_network_demo.py`141142```python143#!/usr/bin/env python3144"""145TP4 Feature Learning - Infinite-Width Network Demo146147Demonstrates usage of the TP4MAML inf module:148- InfGP1LP, InfNTK1LP for infinite-width 1-hidden-layer perceptrons149- InfSGD optimizer150- Utility functions (safe_sqrt, safe_acos, F00ReLU)151152Requires: torch, numpy153Run from the repo root: python scripts/inf_network_demo.py154"""155156import sys157import os158159# Add TP4MAML to path so we can import the inf module160sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..', 'TP4MAML'))161162import torch163import torch.nn as nn164import numpy as np165166167def demo_utils():168 """Demonstrate utility functions from TP4MAML/inf/utils.py"""169 try:170 from inf.utils import safe_sqrt, safe_acos, F00ReLU, MyLinear171 print("=== Utility Functions ===")172173 # safe_sqrt: numerically stable square root174 arr = torch.tensor([-1e-10, 0.0, 1.0, 4.0])175 result = safe_sqrt(arr, eps=1e-6)176 print(f"safe_sqrt({arr.tolist()}) = {result.tolist()}")177178 # safe_acos: numerically stable arccos179 arr2 = torch.tensor([-1.0 - 1e-10, -1.0, 0.0, 1.0, 1.0 + 1e-10])180 result2 = safe_acos(arr2, eps=1e-6)181 print(f"safe_acos(clipped) = {result2.tolist()}")182183 # F00ReLU: ReLU arc-cosine kernel184 # c: cosine similarity, v: variance1, v2: variance2185 c = torch.tensor([0.5])186 v = torch.tensor([1.0])187 v2 = torch.tensor([1.0])188 k = F00ReLU(c, v, v2)189 print(f"F00ReLU(c=0.5, v=1, v2=1) = {k.item():.4f}")190191 # MyLinear: custom linear layer with infinite-width scaling192 linear = MyLinear(in_features=10, out_features=5)193 x = torch.randn(3, 10)194 out = linear(x)195 print(f"MyLinear(10->5) output shape: {out.shape}")196197 except ImportError as e:198 print(f"Could not import inf.utils (run from repo root with TP4MAML in path): {e}")199200201def demo_dynamic_arrays():202 """Demonstrate DynArr and CycArr from TP4MAML/inf/dynamicarray.py"""203 try:204 from inf.dynamicarray import DynArr, CycArr205 print("\n=== Dynamic Arrays ===")206207 # DynArr: growing array for storing activations208 darr = DynArr()209 for i in range(5):210 darr.append(torch.randn(3, 4))211 print(f"DynArr length after 5 appends: {len(darr)}")212213 # CycArr: cyclic (ring buffer) array214 carr = CycArr(capacity=3)215 for i in range(6):216 carr.append(torch.tensor([float(i)]))217 print(f"CycArr (capacity=3) last 3 values: {[carr[i].item() for i in range(len(carr))]}")218219 except ImportError as e:220 print(f"Could not import inf.dynamicarray: {e}")221222223def demo_inf_models():224 """Demonstrate InfGP1LP, FinGP1LP, InfNTK1LP from TP4MAML/inf/inf1lp.py"""225 try:226 from inf.inf1lp import InfGP1LP, FinGP1LP, InfNTK1LP227 print("\n=== Infinite-Width 1-Layer Perceptron Models ===")228229 input_dim = 16230 output_dim = 5231 batch_size = 8232233 x_train = torch.randn(batch_size, input_dim)234 y_train = torch.randint(0, output_dim, (batch_size,))235 x_test = torch.randn(4, input_dim)236237 # InfGP1LP: Infinite-width Gaussian Process limit238 print("Building InfGP1LP...")239 gp_model = InfGP1LP(input_dim=input_dim, output_dim=output_dim)240 print(f" InfGP1LP created: {type(gp_model).__name__}")241242 # InfNTK1LP: Infinite-width NTK limit243 print("Building InfNTK1LP...")244 ntk_model = InfNTK1LP(input_dim=input_dim, output_dim=output_dim)245 print(f" InfNTK1LP created: {type(ntk_model).__name__}")246247 # FinGP1LP: Finite baseline248 print("Building FinGP1LP...")249 fin_model = FinGP1LP(input_dim=input_dim, output_dim=output_dim, width=256)250 print(f" FinGP1LP created: {type(fin_model).__name__}")251252 # Forward pass on finite model253 logits = fin_model(x_test)254 print(f" FinGP1LP forward output shape: {logits.shape}")255256 except ImportError as e:257 print(f"Could not import inf.inf1lp: {e}")258 except Exception as e:259 print(f"Error in inf model demo: {e}")260261262def demo_inf_sgd():263 """Demonstrate InfSGD optimizer from TP4MAML/inf/optim.py"""264 try:265 from inf.optim import InfSGD, InfMultiStepLR266 print("\n=== InfSGD Optimizer ===")267268 # Simple model to optimize269 model = nn.Linear(10, 5)270271 # InfSGD: SGD with infinite-width scaling272 optimizer = InfSGD(model.parameters(), lr=0.01, momentum=0.9)273 print(f"InfSGD created with lr=0.01, momentum=0.9")274275 # InfMultiStepLR scheduler276 scheduler = InfMultiStepLR(optimizer, milestones=[10, 20], gamma=0.1)277 print(f"InfMultiStepLR created with milestones=[10, 20], gamma=0.1")278279 # Simulate a few training steps280 x = torch.randn(4, 10)281 y = torch.randint(0, 5, (4,))282 criterion = nn.CrossEntropyLoss()283284 for step in range(3):285 optimizer.zero_grad()286 out = model(x)287 loss = criterion(out, y)288 loss.backward()289 optimizer.step()290 scheduler.step()291 print(f" Step {step+1}: loss={loss.item():.4f}, lr={scheduler.get_last_lr()}")292293 except ImportError as e:294 print(f"Could not import inf.optim: {e}")295 except Exception as e:296 print(f"Error in InfSGD demo: {e}")297298299def demo_maml_structure():300 """Show the MAML metalearner interface (TP4MAML/meta/maml/metalearners/)"""301 try:302 sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..', 'TP4MAML', 'meta'))303 from maml.metalearners.maml import MAML304 print("\n=== MAML Metalearner ===")305 print(f"MAML class imported: {MAML}")306 print(" MAML implements standard Model-Agnostic Meta-Learning.")307 print(" Use train.py --inf flag to switch to InfMAML.")308 except ImportError as e:309 print(f"\nCould not import MAML metalearner: {e}")310311312if __name__ == "__main__":313 print("TP4 Feature Learning in Infinite-Width Neural Networks - Demo\n")314 demo_utils()315 demo_dynamic_arrays()316 demo_inf_models()317 demo_inf_sgd()318 demo_maml_structure()319 print("\nDemo complete.")320```321322### `scripts/maml_training_demo.py`323324```python325#!/usr/bin/env python3326"""327TP4 MAML Training Demo328329Shows how to invoke MAML training programmatically (mirroring train.py usage).330Demonstrates the dataset loading and metalearner API.331332Requires: torch, torchmeta (see TP4MAML/meta/requirements.txt)333Run from TP4MAML/meta/: python ../../scripts/maml_training_demo.py334"""335336import sys337import os338import argparse339340# Adjust paths for running from repo root341MAML_META_PATH = os.path.join(os.path.dirname(__file__), '..', 'TP4MAML', 'meta')342MAML_INF_PATH = os.path.join(os.path.dirname(__file__), '..', 'TP4MAML')343sys.path.insert(0, MAML_META_PATH)344sys.path.insert(0, MAML_INF_PATH)345346347def build_omniglot_dataset(data_folder: str, num_ways: int = 5, num_shots: int = 1,348 num_shots_test: int = 15):349 """350 Build Omniglot meta-learning dataset using CachedOmniglot.351352 Args:353 data_folder: Path to Omniglot data directory.354 num_ways: Number of classes per episode (N-way).355 num_shots: Number of support examples per class (K-shot).356 num_shots_test: Number of query examples per class.357358 Returns:359 Tuple of (meta_train_dataset, meta_val_dataset, meta_test_dataset)360 """361 try:362 from cached_omniglot import CachedOmniglot363 import torchmeta364 from torchmeta.transforms import ClassSplitter, Categorical365 from torchvision.transforms import Compose, Resize, ToTensor366367 transform = Compose([Resize(28), ToTensor()])368369 meta_train = CachedOmniglot(370 data_folder,371 num_classes_per_task=num_ways,372 transform=transform,373 target_transform=Categorical(num_ways),374 class_augmentations=[torchmeta.transforms.Rotation([90, 180, 270])],375 meta_train=True,376 dataset_transform=ClassSplitter(377 shuffle=True,378 num_support_per_class=num_shots,379 num_query_per_class=num_shots_test380 )381 )382 print(f"CachedOmniglot meta-train: {len(meta_train)} tasks")383 return meta_train384385 except ImportError as e:386 print(f"Could not build Omniglot dataset (missing torchmeta or data): {e}")387 return None388 except Exception as e:389 print(f"Dataset error: {e}")390 return None391392393def demonstrate_metalearner_api():394 """395 Show the API surface of MAML and InfMAML metalearners.396 """397 print("=== MAML / InfMAML API ===\n")398399 try:400 from maml.metalearners.maml import MAML401 print(f"MAML class: {MAML.__module__}.{MAML.__name__}")402 print(f" __init__ signature: see references/api_reference.md")403 print(f" Key methods: train(), evaluate(), get_outer_loss()")404 except ImportError as e:405 print(f"MAML import failed: {e}")406407 try:408 from maml.metalearners.infmaml import InfMAML409 print(f"\nInfMAML class: {InfMAML.__module__}.{InfMAML.__name__}")410 print(f" Extends MAML for infinite-width networks using InfGP1LP/InfNTK1LP")411 except ImportError as e:412 print(f"InfMAML import failed: {e}")413414 try:415 from maml.metalearners.meta_sgd import MetaSGD416 print(f"\nMetaSGD class: {MetaSGD.__module__}.{MetaSGD.__name__}")417 print(f" Meta-SGD variant with per-parameter learned learning rates")418 except ImportError as e:419 print(f"MetaSGD import failed: {e}")420421422def show_training_command_equivalents():423 """424 Print the equivalent train.py CLI commands for common configurations.425 """426 print("\n=== Equivalent train.py Commands ===\n")427428 configs = [429 {430 "description": "5-way 1-shot Omniglot, finite MAML",431 "cmd": "python train.py --dataset omniglot --num-ways 5 --num-shots 1 --num-steps 5"432 },433 {434 "description": "5-way 1-shot Omniglot, infinite-width MAML (GP limit)",435 "cmd": "python train.py --dataset omniglot --num-ways 5 --num-shots 1 --num-steps 5 --inf --use-gp"436 },437 {438 "description": "5-way 5-shot Omniglot, infinite-width MAML (NTK limit)",439 "cmd": "python train.py --dataset omniglot --num-ways 5 --num-shots 5 --num-steps 5 --inf"440 },441 {442 "description": "Run all experiments (uses train_all.sh)",443 "cmd": "cd TP4MAML && bash train_all.sh"444 },445 ]446447 for cfg in configs:448 print(f"# {cfg['description']}")449 print(f" {cfg['cmd']}\n")450451452def show_word2vec_commands():453 """454 Show Word2Vec experiment commands.455 """456 print("=== Word2Vec Experiment Commands ===\n")457458 steps = [459 ("Build C binaries", "cd Word2Vec && make"),460 ("Prepare text8 data", "bash Word2Vec/scripts/create-text8-data.sh"),461 ("Prepare fil9 data", "bash Word2Vec/scripts/create-fil9-data.sh"),462 ("Train text8 (finite)", "bash Word2Vec/scripts/train-text8.sh"),463 ("Train text8 (infinite)", "bash Word2Vec/scripts/train-text8-inf.sh"),464 ("Train fil9 (finite)", "bash Word2Vec/scripts/train-fil9.sh"),465 ("Train fil9 (infinite)", "bash Word2Vec/scripts/train-fil9-inf.sh"),466 ("Evaluate embeddings", "bash Word2Vec/scripts/evaluate.sh"),467 ]468469 for name, cmd in steps:470 print(f"# {name}")471 print(f" {cmd}\n")472473474if __name__ == "__main__":475 print("TP4 MAML Training Demo\n")476 demonstrate_metalearner_api()477 show_training_command_equivalents()478 show_word2vec_commands()479480 # Optionally build dataset if data path provided481 if len(sys.argv) > 1:482 data_folder = sys.argv[1]483 print(f"\n=== Building Omniglot Dataset from {data_folder} ===")484 dataset = build_omniglot_dataset(data_folder, num_ways=5, num_shots=1)485 if dataset is not None:486 print("Dataset built successfully.")487 else:488 print("\nTip: Pass a data folder path as argument to test dataset loading.")489 print(" python maml_training_demo.py /path/to/omniglot/data")490```