TGraphX › Insights › Multimodal Knowledge Graph with Vision and Text Entities: End-to-End Tutorial
← Back to Insights

Multimodal Knowledge Graph with Vision and Text Entities: End-to-End Tutorial

Target keyword: multimodal knowledge graph vision text pytorch

Multimodal Knowledge Graph with Vision and Text Entities: End-to-End Tutorial

Most knowledge graph tutorials work with scalar or low-dimensional entity features — a 64-dimensional embedding per entity, perhaps. Real-world knowledge bases are rarely this clean. An e-commerce product catalog has images. A medical ontology has clinical notes. A scientific database has paper abstracts and figures. When entity representations must combine vision and text, a standard knowledge graph embedding pipeline breaks at the feature ingestion step.

TGraphX's tgraphx.kg module is designed with this problem in mind. It accepts tensor-valued entity features, including image tensors of shape [C, H, W] and variable-length text embeddings as 1D vectors, within the same knowledge graph structure. This tutorial walks through the full pipeline: defining entities with heterogeneous features, building a knowledge graph, training a link prediction model, and evaluating performance.


What Multimodal Knowledge Graphs Add

A standard knowledge graph represents facts as triples (head, relation, tail). TransE, DistMult, and similar embedding models learn a vector for each entity and relation, then score triples by measuring geometric relationships in that embedding space. These methods learn purely from the graph structure and ignore any side information attached to entities.

Multimodal knowledge graph methods augment entity representations with additional modalities. An image-bearing entity gets an initial feature vector derived from a CNN or vision transformer applied to its image. A text-bearing entity gets a feature vector from a sentence encoder applied to its description. These features are then combined with the structural embeddings learned from graph topology.

The benefit is that new entities with images or text but few relational triples can still receive meaningful representations. This is the classic cold-start problem in knowledge bases: a newly added entity with rich side information but no established connections would have a zero or random structural embedding in a pure graph model.


Prerequisites

This tutorial assumes familiarity with:
- PyTorch tensor operations
- Basic knowledge graph concepts (entities, relations, triples)
- The TGraphX Graph object (covered in MNIST as a graph)
- Knowledge graph embedding fundamentals (covered in knowledge graph embedding with TGraphX)

Install TGraphX if you have not already:

bash
pip install tgraphx
        

Data Structure: Heterogeneous Entity Features

In this tutorial we build a small knowledge graph about scientific publications. Each entity is either:
- A paper entity, represented by an image of its first figure (shape [3, 224, 224])
- An author entity, represented by a text embedding of their biography (shape [768], from any sentence encoder)
- A venue entity, represented by a short text embedding of the venue description (shape [384])

python
import torch
        from tgraphx.kg.data import KnowledgeGraphDataset
        
        # --- Entity features ---
        num_papers = 50
        num_authors = 30
        num_venues = 10
        total_entities = num_papers + num_authors + num_venues
        
        # Paper entities: image features [C, H, W]
        paper_image_features = torch.randn(num_papers, 3, 224, 224)
        
        # Author entities: text embeddings [D]
        author_text_features = torch.randn(num_authors, 768)
        
        # Venue entities: shorter text embeddings [D']
        venue_text_features = torch.randn(num_venues, 384)
        
        # Entity type indices (used by the multimodal encoder)
        entity_types = torch.cat([
            torch.zeros(num_papers, dtype=torch.long),    # type 0: paper
            torch.ones(num_authors, dtype=torch.long),     # type 1: author
            torch.full((num_venues,), 2, dtype=torch.long) # type 2: venue
        ])
        

Relation triples connect these entities. Index conventions: paper indices are 0..49, author indices are 50..79, venue indices are 80..89.

python
# --- Relation triples (head, relation, tail) ---
        # Relations: 0 = authored_by, 1 = published_at, 2 = cites
        triples = torch.tensor([
            # paper 0 authored_by author 50
            [0, 0, 50],
            [0, 0, 51],
            [1, 0, 50],
            # paper 0 published_at venue 80
            [0, 1, 80],
            [1, 1, 81],
            # paper 0 cites paper 2
            [0, 2, 2],
            [1, 2, 3],
            [2, 2, 4],
            # ... add more triples for a realistic dataset
        ], dtype=torch.long)
        
        print(f"Triples shape: {triples.shape}")  # [num_triples, 3]
        

Building the Multimodal Knowledge Graph Dataset

The tgraphx.kg module accepts a KnowledgeGraphDataset that holds entities, their features, and the triple list:

python
from tgraphx.kg.data import KnowledgeGraphDataset
        from tgraphx.kg.multimodal import MultimodalEntityEncoder
        
        # The dataset needs a flat entity feature dictionary
        # mapping entity type -> feature tensor
        entity_features = {
            "paper": paper_image_features,     # [50, 3, 224, 224]
            "author": author_text_features,    # [30, 768]
            "venue": venue_text_features,      # [10, 384]
        }
        
        kg_dataset = KnowledgeGraphDataset(
            triples=triples,
            num_entities=total_entities,
            num_relations=3,
            entity_types=entity_types,
            entity_features=entity_features,
        )
        
        print(f"Entities: {kg_dataset.num_entities}")
        print(f"Relations: {kg_dataset.num_relations}")
        print(f"Triples: {len(kg_dataset)}")
        

The Multimodal Entity Encoder

The encoder projects each modality into a shared embedding space. Images go through a small convolutional encoder; text vectors go through a linear projection. The output is a uniform embedding of fixed size embed_dim for every entity, regardless of its original modality:

python
from tgraphx.kg.multimodal import MultimodalEntityEncoder
        
        encoder = MultimodalEntityEncoder(
            embed_dim=128,
            modality_configs={
                "paper": {
                    "type": "image",
                    "in_channels": 3,
                    "spatial_size": 224,
                },
                "author": {
                    "type": "vector",
                    "in_dim": 768,
                },
                "venue": {
                    "type": "vector",
                    "in_dim": 384,
                },
            }
        )
        
        # Encode all entities into [total_entities, embed_dim]
        # Each entity type is encoded separately then concatenated
        paper_embs = encoder.encode_modality("paper", paper_image_features)  # [50, 128]
        author_embs = encoder.encode_modality("author", author_text_features) # [30, 128]
        venue_embs = encoder.encode_modality("venue", venue_text_features)    # [10, 128]
        
        all_entity_embeddings = torch.cat([paper_embs, author_embs, venue_embs], dim=0)
        print(all_entity_embeddings.shape)  # [90, 128]
        

Link Prediction Model

With entity embeddings in a shared space, we can apply any standard scoring function for link prediction. TGraphX's tgraphx.kg.models module provides TransE and DistMult implementations that operate on pre-computed embeddings:

python
from tgraphx.kg.models import DistMultScorer
        import torch.nn as nn
        
        class MultimodalKGModel(nn.Module):
            def __init__(self, encoder, num_relations, embed_dim):
                super().__init__()
                self.encoder = encoder
                # Relation embeddings are learned, entity embeddings come from encoder
                self.relation_emb = nn.Embedding(num_relations, embed_dim)
                self.scorer = DistMultScorer()
        
            def forward(self, entity_features_by_type, triples):
                # Encode all entities
                paper_embs = self.encoder.encode_modality(
                    "paper", entity_features_by_type["paper"]
                )
                author_embs = self.encoder.encode_modality(
                    "author", entity_features_by_type["author"]
                )
                venue_embs = self.encoder.encode_modality(
                    "venue", entity_features_by_type["venue"]
                )
                all_embs = torch.cat([paper_embs, author_embs, venue_embs], dim=0)
        
                heads = all_embs[triples[:, 0]]    # [B, D]
                rels = self.relation_emb(triples[:, 1])  # [B, D]
                tails = all_embs[triples[:, 2]]    # [B, D]
        
                return self.scorer(heads, rels, tails)
        
        model = MultimodalKGModel(encoder, num_relations=3, embed_dim=128)
        

Training Loop

python
import torch.optim as optim
        from tgraphx.kg.trainer import KGTrainer
        
        # Use TGraphX trainer for standardized KG training
        trainer = KGTrainer(
            model=model,
            dataset=kg_dataset,
            optimizer=optim.Adam(model.parameters(), lr=1e-3),
            loss="margin_ranking",   # standard KG training objective
            neg_samples=10,          # negatives per positive triple
            epochs=100,
            device="cpu",
        )
        
        history = trainer.fit(entity_features_by_type=entity_features)
        
        print(f"Final MRR: {history['mrr'][-1]:.4f}")
        print(f"Final Hits@10: {history['hits_at_10'][-1]:.4f}")
        

Evaluation

The tgraphx.kg.evaluation module provides standard KG evaluation metrics:

python
from tgraphx.kg.evaluation import evaluate_link_prediction
        
        results = evaluate_link_prediction(
            model=model,
            test_triples=kg_dataset.test_triples,
            all_triples=kg_dataset.all_triples,
            entity_features=entity_features,
            num_entities=total_entities,
            metrics=["mrr", "hits_at_1", "hits_at_3", "hits_at_10"],
            filtered=True,  # filter out true triples from ranking (standard protocol)
        )
        
        for metric, value in results.items():
            print(f"{metric}: {value:.4f}")
        

The filtered=True flag enables filtered evaluation, which is standard in the KG literature: when ranking candidate tails for a query (head, relation, ?), true tails from the training set are excluded so the model is not penalized for assigning them high scores.


Handling Missing Modalities

In real datasets, not all entities will have all modalities. A paper might lack a figure image; an author might lack a biography text. A practical pattern for missing modalities is to replace them with a learned type-specific default embedding:

python
# Replace missing paper features with a learned "unknown" embedding
        unknown_paper_emb = nn.Parameter(torch.randn(1, 3, 224, 224))
        
        for i, paper_id in enumerate(paper_ids_with_missing_images):
            paper_image_features[paper_id] = unknown_paper_emb.squeeze(0).detach()
        

This is a simple approximation. More sophisticated approaches include masking losses from missing-modality entities during training or using modality-dropout regularization.


Limitations and Honest Notes

Image encoding is expensive. Processing 50,000 entities with [3, 224, 224] image features requires either a GPU with sufficient memory or pre-computing and caching embeddings. The encoder is not designed to handle millions of image entities in a single forward pass.

Text embeddings must be pre-computed externally. TGraphX does not include a text encoder (BERT, sentence-transformers, etc.). The tutorial assumes you have already converted text to fixed-size vectors using an external library. This is by design — text encoding is a rapidly evolving area and coupling it to a graph library would create an unnecessary dependency.

DistMult is a symmetric scorer. It cannot model asymmetric relations like "authored_by" vs "is_author_of" correctly. For asymmetric relations, use TransE or a relation-specific MLP head.

The cold-start benefit requires sufficient feature quality. If the image or text encoder is poorly pre-trained, the multimodal embeddings will not be informative, and the model will underperform a purely structural approach.

Small graphs expose the limitation of structural KG methods. With only 90 entities and a few hundred triples, there is insufficient data for learned embeddings to converge well. These methods are designed for knowledge bases with thousands to millions of entities.

For related work on tensor-valued features in graph neural networks, see the graph generation with tensor-valued node features article. For knowledge graph embedding fundamentals in TGraphX, see the knowledge graph embedding tutorial.


Negative Sampling for KG Training

Knowledge graph training requires negative examples — entity pairs that are not connected by the given relation. Standard KG training samples random negatives by replacing either the head or tail entity with a random entity from the vocabulary:

python
from tgraphx.kg.sampling import NegativeSampler
        
        neg_sampler = NegativeSampler(
            num_entities=total_entities,
            num_negatives=5,    # 5 negatives per positive triple
            strategy="random",  # or "type_constrained" if entity types matter
        )
        
        # During training
        for batch_triples in kg_dataset.train_loader(batch_size=64):
            pos_heads = batch_triples[:, 0]
            pos_rels = batch_triples[:, 1]
            pos_tails = batch_triples[:, 2]
        
            neg_heads, neg_tails = neg_sampler.sample(pos_heads, pos_rels, pos_tails)
            # neg_heads/tails: [B * num_negatives] — corrupt head or tail randomly
        

For type-constrained sampling, the sampler only replaces entities with others of the same type. In a knowledge graph where authors are type 1 and papers are type 0, corrupting an "authored_by" triple's tail would only generate author entity negatives. This produces harder, more meaningful negatives for multimodal graphs where entity types carry strong semantic meaning.


Hyperparameter Sensitivity

Multimodal KG models have more hyperparameters than standard KG models because the modality encoders add their own parameters. Key hyperparameters to tune:

Embedding dimension. A value between 64 and 256 is typical. Too small and the model cannot represent entity relationships; too large and it overfits on small graphs. Start with 128.

Learning rate schedule. A cosine annealing schedule with a warmup period tends to work better than fixed learning rates for multimodal models, because the image encoder parameters need lower learning rates than the scoring function.

Modality weight. Some implementations include a learned weight balancing the structural embedding versus the modality embedding for each entity. This is especially important when not all entities have the same modalities.

Negative sampling ratio. Higher ratios (10–20 negatives per positive) tend to improve ranking metrics at the cost of training time. For very small graphs, lower ratios (3–5) are sufficient.


Extending to New Modalities

The MultimodalEntityEncoder architecture in TGraphX is designed to be extensible. Any modality that can be projected to a fixed-size embedding can be added:

python
# Adding an audio modality (e.g., 1D waveform features)
        encoder = MultimodalEntityEncoder(
            embed_dim=128,
            modality_configs={
                "paper": {"type": "image", "in_channels": 3, "spatial_size": 224},
                "author": {"type": "vector", "in_dim": 768},
                "venue": {"type": "vector", "in_dim": 384},
                "talk": {
                    "type": "vector",
                    "in_dim": 1024,  # audio segment embedding from an external encoder
                },
            }
        )
        

The "vector" type handles any fixed-size input, so any modality that has been pre-encoded into a fixed-size vector can be added. Modalities with variable structure (e.g., point clouds, graphs-of-graphs) would require custom encoder modules.


What This Tutorial Builds On

This tutorial extends the concepts in the knowledge graph embedding tutorial by adding tensor-valued entity features. If you are new to TGraphX's KG module, start with that article to understand the basic scoring functions and evaluation protocol before attempting this multimodal extension. The MNIST as a graph tutorial provides background on tensor-valued node features in the simpler node classification setting.

See the articles hub for the full index of TGraphX guides.