rohitg00/ai-engineering-from-scratch · error · TypeError

token_embedding must be a TokenEmbedding

Error message

token_embedding must be a TokenEmbedding

What it means

Error "token_embedding must be a TokenEmbedding" thrown in rohitg00/ai-engineering-from-scratch.

Source

Thrown at phases/19-capstone-projects/32-token-positional-embeddings/code/main.py:139

                f"seq_len {seq_len} exceeds max_context_length {self.max_context_length}"
            )
        return self.pe[:seq_len]


class EmbeddingComposer(nn.Module):
    """Sums a token embedding with a positional embedding.

    The positional embedding may be learned or sinusoidal.
    """

    def __init__(
        self,
        token_embedding: TokenEmbedding,
        positional_embedding: nn.Module,
    ) -> None:
        super().__init__()
        if not isinstance(token_embedding, TokenEmbedding):
            raise TypeError("token_embedding must be a TokenEmbedding")
        if not isinstance(
            positional_embedding,
            (LearnedPositionalEmbedding, SinusoidalPositionalEmbedding),
        ):
            raise TypeError(
                "positional_embedding must be Learned or Sinusoidal Positional Embedding"
            )
        if token_embedding.d_model != getattr(positional_embedding, "d_model", None):
            raise ValueError("token and positional embeddings must share d_model")
        self.token_embedding = token_embedding
        self.positional_embedding = positional_embedding

    @property
    def d_model(self) -> int:
        return self.token_embedding.d_model

    def forward(self, ids: torch.Tensor) -> torch.Tensor:
        if ids.dim() != 2:

View on GitHub (pinned to 39ea8a1c6d)

When it happens

Trigger: Thrown at phases/19-capstone-projects/32-token-positional-embeddings/code/main.py:139 when the library encounters an invalid state.

Common situations: See trigger scenarios.


AI-assisted analysis of rohitg00/ai-engineering-from-scratch@39ea8a1c6d (2026-08-26). Data as JSON: /api/errors/d0f3ffd39d704001. Report an issue: GitHub.