Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions doc/modules/high_level_api.rst
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@ Models

ConceptBottleneckModel
ConceptEmbeddingModel
ConceptMemoryReasoner
GraphConceptBottleneckModel
CausallyReliableConceptBottleneckModel
BlackBox
Expand Down
1 change: 0 additions & 1 deletion examples/contributing/model.md
Original file line number Diff line number Diff line change
Expand Up @@ -393,7 +393,6 @@ from torch_concepts.nn import (
#### Special Layers
```python
from torch_concepts.nn import (
SelectorLatentToExogenous, # Memory-augmented selection
WANDAGraphLearner, # Learn concept graph structure
)
```
Expand Down
115 changes: 115 additions & 0 deletions examples/utilization/0_layer/7_concept_based_memory_reasoner.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,115 @@
"""
Example: Concept Memory Reasoner with Low-Level API

This example demonstrates how to build a Concept Memory Reasoner (CMR)
using the low-level encoder and predictor layers.
"""
import torch
from sklearn.metrics import accuracy_score
from torch.nn import ModuleDict

from torch_concepts import seed_everything
from torch_concepts.data.datasets import ToyDataset
from torch_concepts.nn import (
RuleMemory,
RuleReconstructionPredictor,
RuleTaskPredictor,
CategoricalSelector,
LinearEmbeddingToConcept,
)


def main():
latent_dims = 10
n_epochs = 500
n_samples = 1000
nb_rules = 10
memory_latent_size = 100

seed_everything(42)

dataset = ToyDataset(dataset='xor', seed=42, n_gen=n_samples)
x_train = dataset.input_data
concept_idx = list(dataset.graph.edge_index[0].unique().numpy())
task_idx = list(dataset.graph.edge_index[1].unique().numpy())
c_train = dataset.concepts[:, concept_idx]
y_train = dataset.concepts[:, task_idx]

n_features = x_train.shape[1]
n_concepts = c_train.shape[1]
n_tasks = y_train.shape[1]

latent_encoder = torch.nn.Sequential(
torch.nn.Linear(n_features, latent_dims),
torch.nn.LeakyReLU(),
)
selector_encoder = CategoricalSelector(
in_latent=latent_dims,
out_concepts=n_tasks,
out_exogenous=nb_rules,
)
concept_encoder = LinearEmbeddingToConcept(in_embeddings=latent_dims, out_concepts=n_concepts)
memory = RuleMemory(
n_tasks=n_tasks,
n_rules=nb_rules,
n_concepts=n_concepts,
latent_size=memory_latent_size,
)
task_predictor = RuleTaskPredictor(
in_concepts=n_concepts,
in_exogenous=nb_rules,
out_concepts=n_tasks,
)
reconstruction_predictor = RuleReconstructionPredictor(
in_concepts=n_concepts,
in_exogenous=nb_rules,
out_concepts=n_tasks,
rec_weight=0.1,
)

model = ModuleDict({
'latent_encoder': latent_encoder,
'selector_encoder': selector_encoder,
'concept_encoder': concept_encoder,
'memory': memory,
'task_predictor': task_predictor,
'reconstruction_predictor': reconstruction_predictor,
})

optimizer = torch.optim.AdamW(model.parameters(), lr=0.01)
concept_loss_fn = torch.nn.BCEWithLogitsLoss()
task_loss_fn = torch.nn.BCELoss(reduction='none')
model.train()

for epoch in range(n_epochs):
optimizer.zero_grad()

emb = latent_encoder(x_train)
selector = selector_encoder(latent=emb)
c_logits = concept_encoder(embeddings=emb)
c_probs = c_logits.sigmoid()
roles = memory()

y_pred = task_predictor(concepts=c_probs, selector=selector, roles=roles)
y_pred_with_rec = reconstruction_predictor(concepts=c_probs, selector=selector, roles=roles)

concept_loss = concept_loss_fn(c_logits, c_train)
task_loss_no_rec = task_loss_fn(y_pred, y_train)
task_loss_with_rec = task_loss_fn(y_pred_with_rec, y_train)
switched_task_loss = ((1.0 - y_train) * task_loss_no_rec + y_train * task_loss_with_rec).mean()
loss = concept_loss + switched_task_loss

loss.backward()
optimizer.step()

if epoch % 100 == 0:
task_accuracy = accuracy_score(y_train.cpu(), (y_pred.detach() > 0.5).cpu())
concept_accuracy = accuracy_score(c_train.cpu(), (c_logits.detach() > 0.0).cpu())
print(
f'Epoch {epoch}: Loss {loss.item():.2f} | '
f'Task Acc: {task_accuracy:.2f} | Concept Acc: {concept_accuracy:.2f}'
)


if __name__ == '__main__':
main()
56 changes: 49 additions & 7 deletions examples/utilization/2.2_model/10_different_training_modes.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,9 +22,10 @@

import torch
from torch_concepts import seed_everything
from torch_concepts.nn import ConceptBottleneckModel, ConceptEmbeddingModel, MLP
from torch_concepts.nn import ConceptBottleneckModel, ConceptEmbeddingModel, MLP, ConceptMemoryReasoner, CMRBlendedLoss, ConceptLoss
from torch_concepts.nn import DeterministicInference, IndependentInference
from torch_concepts.data import ToyDataset
from torch_concepts.data.datasets import ToyDataset
from torch_concepts.data.base.datamodule import ConceptDataModule
from torch.distributions import Bernoulli

Expand All @@ -34,7 +35,7 @@


def evaluate(model, datamodule, n_concepts, query):
"""Evaluate model on test set and return concept/task accuracy."""
"""Evaluate model on a data split and return concept/task accuracy."""
concept_acc_fn = BinaryAccuracy()
task_acc_fn = BinaryAccuracy()

Expand All @@ -48,10 +49,12 @@ def evaluate(model, datamodule, n_concepts, query):
for batch in test_loader:
# model.eval() automatically selects eval_inference
out = model(input=batch['inputs']['x'], query=query)
c_logits = out.logits[:, :n_concepts]
y_logits = out.logits[:, n_concepts:]
c_pred = torch.sigmoid(c_logits)
y_pred = torch.sigmoid(y_logits)
predictions = out.logits if out.logits is not None else out.probs
c_pred = predictions[:, :n_concepts]
y_pred = predictions[:, n_concepts:]
if out.logits is not None:
c_pred = torch.sigmoid(c_pred)
y_pred = torch.sigmoid(y_pred)

c_true = batch['concepts']['c'][:, :n_concepts]
y_true = batch['concepts']['c'][:, n_concepts:]
Expand Down Expand Up @@ -104,7 +107,7 @@ def main():

# Define variable distributions as Bernoulli
variable_distributions = {name: Bernoulli for name in concept_names}
loss = torch.nn.BCEWithLogitsLoss()
loss = ConceptLoss(binary=torch.nn.BCEWithLogitsLoss(), binary_param="logits")
optim = torch.optim.AdamW
optim_kwargs = {'lr': 0.1}

Expand Down Expand Up @@ -202,6 +205,45 @@ def main():
trainer_cem.fit(model_cem, datamodule=datamodule)
evaluate(model_cem, datamodule, n_concepts, query)

# =========================================================================
# CMR WITH JOINT TRAINING
# =========================================================================
print("\n" + "=" * 60)
print("Example 4: CMR with Joint Training")
print("=" * 60)
print("Uses DeterministicInference for both training and evaluation")

cmr_loss = CMRBlendedLoss(task_names=['xor'])
optim_kwargs_cmr = {'lr': 0.01}

model_cmr = ConceptMemoryReasoner(
input_size=n_features,
annotations=annotations,
backbone=MLP(input_size=n_features, hidden_size=16, n_layers=1),
latent_size=16,
variable_distributions=variable_distributions,
task_names=['xor'],
n_rules=10,
memory_latent_size=100,
memory_decoder_hidden_layers=1,
selector_hidden_layers=1,
hard_roles_at_eval=True,
inference=DeterministicInference,
train_inference=DeterministicInference,
lightning=True,
loss=cmr_loss,
rec_weight=0,
optim_class=optim,
optim_kwargs=optim_kwargs_cmr,
)
print(f"Model type: {type(model_cmr).__name__}")
print(f"Eval inference: {model_cmr.eval_inference.__class__.__name__}")
print(f"Training inference: {model_cmr.train_inference.__class__.__name__}")

trainer_cmr = Trainer(max_epochs=100)
trainer_cmr.fit(model_cmr, datamodule=datamodule)
evaluate(model_cmr, datamodule, n_concepts, query)


if __name__ == "__main__":
main()
26 changes: 26 additions & 0 deletions tests/nn/modules/high/models/test_cmr.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,26 @@
import torch
from torch_concepts import Annotations
from torch_concepts.nn import CMRBlendedLoss
from torch_concepts.nn.modules.high.models.cmr import ConceptMemoryReasoner


def test_cmr_routes_reconstruction_prediction_through_modeloutput_extra():
model = ConceptMemoryReasoner(
input_size=2,
annotations=Annotations(labels=["c1", "c2", "xor"], cardinalities=[1, 1, 1]),
task_names=["xor"],
n_rules=3,
)
target = torch.tensor([[0., 1., 1.], [1., 0., 0.]])
query = model.build_query(target)
query["tasks_with_rec"] = None
output = model(query=query, evidence={"input": torch.randn(2, 2)})

assert output.probs["xor"].shape == (2, 1)
assert output.extra["task_input"].shape == (2, 1)
assert output.extra["input_with_rec"].shape == (2, 1)
assert "tasks_with_rec" not in output.probs.annotation.label_to_index

loss = CMRBlendedLoss(task_names=["xor"])(output, model.prepare_target(target))
loss.backward()
assert torch.isfinite(loss)
107 changes: 107 additions & 0 deletions tests/nn/modules/low/encoders/test_selector.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import torch
import torch.nn as nn
from torch_concepts.nn.modules.low.dense_layers import SelectorEmbeddingEncoder
from torch_concepts.nn.modules.low.encoders.selector import CategoricalSelector


class TestSelectorEmbeddingEncoder(unittest.TestCase):
Expand Down Expand Up @@ -119,5 +120,111 @@ def test_batch_processing(self):
self.assertEqual(output.shape, (batch_size, 3, 4))


class TestCategoricalSelector(unittest.TestCase):
"""Test CategoricalSelector."""

def test_initialization(self):
"""Test selector initialization."""
selector = CategoricalSelector(
in_latent=64,
out_concepts=5,
out_exogenous=8,
selector_hidden_layers=2,
)
self.assertEqual(selector.in_latent, 64)
self.assertEqual(selector.out_concepts, 5)
self.assertEqual(selector.out_exogenous, 8)
self.assertEqual(selector.selector_hidden_layers, 2)

def test_forward_shape(self):
"""Test forward pass output shape."""
selector = CategoricalSelector(
in_latent=64,
out_concepts=4,
out_exogenous=6,
)
latent = torch.randn(2, 64)
output = selector(latent=latent)
self.assertEqual(output.shape, (2, 4, 6))

def test_output_is_normalized_over_exogenous_dim(self):
"""Test output probabilities sum to 1 over exogenous dimension."""
selector = CategoricalSelector(
in_latent=32,
out_concepts=3,
out_exogenous=5,
)
latent = torch.randn(3, 32)
output = selector(latent=latent)

sums = output.sum(dim=-1)
self.assertTrue(torch.allclose(sums, torch.ones_like(sums), atol=1e-5))

def test_gradient_flow(self):
"""Test gradient flow through selector."""
selector = CategoricalSelector(
in_latent=32,
out_concepts=3,
out_exogenous=4,
)
embeddings = torch.randn(2, 32, requires_grad=True)
output = selector(latent=embeddings)
loss = output.sum()
loss.backward()
self.assertIsNotNone(embeddings.grad)

def test_hidden_layers_configuration(self):
"""Test configurable hidden layers in selector network."""
selector_zero = CategoricalSelector(
in_latent=32,
out_concepts=3,
out_exogenous=4,
selector_hidden_layers=0,
)
selector_two = CategoricalSelector(
in_latent=32,
out_concepts=3,
out_exogenous=4,
selector_hidden_layers=2,
)

linear_zero = sum(isinstance(layer, nn.Linear) for layer in selector_zero.selector)
linear_two = sum(isinstance(layer, nn.Linear) for layer in selector_two.selector)

self.assertEqual(linear_zero, 1)
self.assertEqual(linear_two, 3)

def test_selector_hidden_layers_validation(self):
"""Test hidden layer argument validation."""
with self.assertRaises(ValueError):
CategoricalSelector(
in_latent=32,
out_concepts=3,
out_exogenous=4,
selector_hidden_layers=-1,
)

def test_selector_network(self):
"""Test selector network structure."""
selector = CategoricalSelector(
in_latent=64,
out_concepts=4,
out_exogenous=6,
)
self.assertIsInstance(selector.selector, nn.Sequential)

def test_batch_processing(self):
"""Test different batch sizes."""
selector = CategoricalSelector(
in_latent=32,
out_concepts=3,
out_exogenous=4,
)
for batch_size in [1, 4, 8]:
embeddings = torch.randn(batch_size, 32)
output = selector(latent=embeddings)
self.assertEqual(output.shape, (batch_size, 3, 4))


if __name__ == '__main__':
unittest.main()
Loading