huggingface/pytorch-image-models · error · ValueError

Unsupported position embedding shape in {checkpoint_path}: {

Error message

Unsupported position embedding shape in {checkpoint_path}: {pos_embed_w.shape}

What it means

After normalization, the position embedding must be 4D (1, H, W, C) — either reshaped from tokens or already stored HWC (Big Vision NaFlex style). Any other rank (e.g. a 3D or 5D tensor) is unsupported by this loader path.

Source

Thrown at timm/models/vision_transformer.py:1569

                    num_pos_tokens = pos_embed_w.shape[1]
                    grid_size = int(math.sqrt(num_pos_tokens))
                    if grid_size * grid_size != num_pos_tokens:
                        checkpoint_prefix_tokens = (
                            1 if f'{prefix}cls' in w else getattr(embeds, 'num_prefix_tokens', 0))
                        num_pos_tokens -= checkpoint_prefix_tokens
                        grid_size = int(math.sqrt(num_pos_tokens))
                        if grid_size * grid_size != num_pos_tokens:
                            raise ValueError(
                                f'Cannot infer position grid from {pos_embed_w.shape[1]} tokens '
                                f'in {checkpoint_path}')
                        prefix_pos_embed = pos_embed_w[:, :checkpoint_prefix_tokens]
                        pos_embed_w = pos_embed_w[:, checkpoint_prefix_tokens:]
                    pos_embed_w = pos_embed_w.reshape(1, grid_size, grid_size, pos_embed_w.shape[-1])
                else:
                    # Big Vision NaFlex stores the grid directly as HWC.
                    pos_embed_w = pos_embed_w.unsqueeze(0)
            if pos_embed_w.ndim != 4:
                raise ValueError(f'Unsupported position embedding shape in {checkpoint_path}: {pos_embed_w.shape}')

            if prefix_pos_embed is not None:
                prefix_index = 0
                if embeds.cls_token is not None and prefix_pos_embed.shape[1] > prefix_index:
                    embeds.cls_token.add_(prefix_pos_embed[:, prefix_index:prefix_index + 1])
                    prefix_index += 1
                if embeds.reg_token is not None and prefix_pos_embed.shape[1] > prefix_index:
                    num_reg_tokens = min(embeds.reg_token.shape[1], prefix_pos_embed.shape[1] - prefix_index)
                    embeds.reg_token[:, :num_reg_tokens].add_(
                        prefix_pos_embed[:, prefix_index:prefix_index + num_reg_tokens])

            if pos_embed_w.shape != embeds.pos_embed.shape:
                pos_embed_w = resample_abs_pos_embed_nhwc(
                    pos_embed_w,
                    new_size=embeds.pos_embed.shape[1:3],
                    interpolation=interpolation,
                    antialias=antialias,
                    verbose=True,

View on GitHub (pinned to 9a5261e31b)

Solutions

  1. Pre-convert the pos_embed tensor in the checkpoint to shape (1, H, W, C) or flat tokens for a square grid
  2. Use the standard PyTorch loader (checkpoint_filter / torch format) for non-standard checkpoints

Example fix

# before
load_pretrained(model, 'weird_pos.npz')
# after
w['pos_embed'] = w['pos_embed'].reshape(1, H, W, C)  # normalize to HWC
load_pretrained(model, 'weird_pos.npz')
Defensive patterns

Strategy: validation

Validate before calling

assert pos_embed_w.ndim == 2 or pos_embed_w.ndim == 4, f'bad pos embed rank {pos_embed_w.ndim}'

Try / catch

try:
    load_pretrained(model, path)
except ValueError as e:
    if 'Unsupported position embedding shape' in str(e):
        w['pos_embed'] = w['pos_embed'].reshape(1, H, W, C)
        load_pretrained(model, path)
    else:
        raise

Prevention

When it happens

Trigger: A checkpoint whose pos_embed tensor has an exotic layout — e.g. factorized (two tensors), extra batch dims, or a per-head layout — reaching the ndim != 4 check.

Common situations: Loading experimental or community-converted checkpoints that store position embeddings non-standardly.

Related errors


AI-assisted analysis of huggingface/pytorch-image-models@9a5261e31b (2026-08-27). Data as JSON: /api/errors/4dd6329e34f7e5a8. Report an issue: GitHub.