{"record":{"id":"a4df4861f8f80890","repo":"huggingface/pytorch-image-models","slug":"cannot-infer-position-grid-from-pos-embed-w-shape","errorCode":null,"errorMessage":"Cannot infer position grid from {pos_embed_w.shape[1]} tokens in {checkpoint_path}","messagePattern":"Cannot infer position grid from (.+?) tokens in (.+?)","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"timm/models/vision_transformer.py","lineNumber":1559,"sourceCode":"            f'{prefix}pos_embedding' if big_vision\n            else f'{prefix}Transformer/posembed_input/pos_embedding')\n        if embeds.pos_embed is not None and pos_embed_key in w:\n            pos_embed_w = _n2p(w[pos_embed_key], t=False)\n            prefix_pos_embed = None\n            if pos_embed_w.ndim == 2:\n                pos_embed_w = pos_embed_w.unsqueeze(0)\n            if pos_embed_w.ndim == 3:\n                if pos_embed_w.shape[0] == 1:\n                    # Flattened NLC tables may include class/register positions.\n                    num_pos_tokens = pos_embed_w.shape[1]\n                    grid_size = int(math.sqrt(num_pos_tokens))\n                    if grid_size * grid_size != num_pos_tokens:\n                        checkpoint_prefix_tokens = (\n                            1 if f'{prefix}cls' in w else getattr(embeds, 'num_prefix_tokens', 0))\n                        num_pos_tokens -= checkpoint_prefix_tokens\n                        grid_size = int(math.sqrt(num_pos_tokens))\n                        if grid_size * grid_size != num_pos_tokens:\n                            raise ValueError(\n                                f'Cannot infer position grid from {pos_embed_w.shape[1]} tokens '\n                                f'in {checkpoint_path}')\n                        prefix_pos_embed = pos_embed_w[:, :checkpoint_prefix_tokens]\n                        pos_embed_w = pos_embed_w[:, checkpoint_prefix_tokens:]\n                    pos_embed_w = pos_embed_w.reshape(1, grid_size, grid_size, pos_embed_w.shape[-1])\n                else:\n                    # Big Vision NaFlex stores the grid directly as HWC.\n                    pos_embed_w = pos_embed_w.unsqueeze(0)\n            if pos_embed_w.ndim != 4:\n                raise ValueError(f'Unsupported position embedding shape in {checkpoint_path}: {pos_embed_w.shape}')\n\n            if prefix_pos_embed is not None:\n                prefix_index = 0\n                if embeds.cls_token is not None and prefix_pos_embed.shape[1] > prefix_index:\n                    embeds.cls_token.add_(prefix_pos_embed[:, prefix_index:prefix_index + 1])\n                    prefix_index += 1\n                if embeds.reg_token is not None and prefix_pos_embed.shape[1] > prefix_index:\n                    num_reg_tokens = min(embeds.reg_token.shape[1], prefix_pos_embed.shape[1] - prefix_index)","sourceCodeStart":1541,"sourceCodeEnd":1577,"githubUrl":"https://github.com/huggingface/pytorch-image-models/blob/9a5261e31b3b5128526eb2658333b4c0a54464ae/timm/models/vision_transformer.py#L1541-L1577","documentation":"When loading position embeddings stored as a flat token sequence (2D), the loader infers a square grid via sqrt of token count. If neither the raw count nor the count minus prefix tokens is a perfect square, the grid cannot be inferred.","triggerScenarios":"Loading a JAX checkpoint trained on non-square grids (rectangular img_size or non-square patches) — token count like 240 (16x15) fails the square test.","commonSituations":"Loading Big Vision checkpoints fine-tuned at rectangular resolutions, or NaFlex-adjacent checkpoints, into the square-grid ViT loader.","solutions":["Load the checkpoint into a model whose grid matches the training resolution (set img_size/patch_size so the grid is rectangular where supported)","Use a PyTorch-format port of the checkpoint that stores the grid directly (HWC) instead of flat tokens","Skip position embedding loading and retrain/interpolate manually"],"exampleFix":"# before\nmodel = vit_base_patch16_224()\nload_pretrained(model, 'bigvision_rect.npz')\n# after\nmodel = vit_base_patch16_224(img_size=(240, 256))  # match training grid\nload_pretrained(model, 'bigvision_rect.npz')","handlingStrategy":"validation","validationCode":"import math\nn = pos_embed_w.shape[1]\ng = int(math.sqrt(n))\nok = g * g == n or (lambda m: (s := int(math.sqrt(m))) * s == m)(n - (1 if 'cls' in w else 0))\nassert ok, 'non-square grid; use matching img_size'","typeGuard":null,"tryCatchPattern":"try:\n    load_pretrained(model, path)\nexcept ValueError as e:\n    if 'Cannot infer position grid' in str(e):\n        model = rebuild_with_rect_grid(path)  # match training resolution\n        load_pretrained(model, path)\n    else:\n        raise","preventionTips":["Check training resolution of the checkpoint","Compute expected grid = (img_size/patch_size)^2 before loading"],"tags":["timm","vit","checkpoint","pos-embed"],"backgroundTag":"checkpoint-weight-mismatch","analyzedSha":"9a5261e31b3b5128526eb2658333b4c0a54464ae","analyzedAt":"2026-08-27T02:34:25.417Z","schemaVersion":2},"datasetVersion":"2026-08-27T03:17:27.898Z"}