{"record":{"id":"06f48d14b78a41b3","repo":"jax-ml/jax","slug":"can-only-load-scalars-from-smem","errorCode":null,"errorMessage":"Can only load scalars from SMEM","messagePattern":"Can only load scalars from SMEM","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic/lowering.py","lineNumber":2337,"sourceCode":"    physical_out_dtype = physical_element_aval.dtype\n    physical_out_shape = jax_core.physical_shape(\n        aval_out.shape, aval_out.dtype\n    )\n  else:\n    physical_out_dtype = aval_out.dtype\n    physical_out_shape = aval_out.shape\n  if not is_smem_load and not ref_block_shape:\n    raise NotImplementedError(\n        \"Indexing into a ()-shaped Ref not yet supported on TPU.\")\n  starts, sizes, strides, _, _ = _indexer_to_start_size_stride(\n      idx,\n      ref_block_shape,\n      cast_to_index=True,\n  )\n  need_stride = not all((s is None or s == 1) for s in strides)\n  if is_smem_load:\n    if ctx.avals_out[0].shape:\n      raise ValueError(\"Can only load scalars from SMEM\")\n    return _maybe_cast_load_to_bool(ctx, aval_out, memref.load(ref, starts))\n  elif str(ref_type.memory_space) != \"#tpu.memory_space<vmem>\":\n    extra = \"\"\n    if str(ref_type.memory_space) == \"#tpu.memory_space<any>\":\n      extra = \" ANY memory space can only be accessed using async_copy.\"\n    raise ValueError(\n        \"Loads are only allowed on VMEM and SMEM references.\" + extra\n    )\n  load_aval = jax_core.ShapedArray(sizes, dtype=physical_out_dtype)\n  if need_stride:\n    load_val = tpu.strided_load(\n        ctx.aval_to_ir_type(load_aval, is_kernel_boundary=True),\n        ref,\n        starts,\n        strides,\n    )\n  else:\n    load_val = vector.load(","sourceCodeStart":2319,"sourceCodeEnd":2355,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic/lowering.py#L2319-L2355","documentation":"Raised when loading from an SMEM (scalar memory) ref with a non-scalar (array-shaped) output. SMEM on TPU is scalar memory; loads from it must produce scalars. Attempting to load a vector/array block from SMEM is a shape/memory-space mismatch.","triggerScenarios":"pl.load where the ref's memory_space is SMEM but the output aval (or the requested slice) has shape with any non-1 dimensions.","commonSituations":"Setting memory_space=SMEM on a BlockSpec for a tensor (array) input to try to 'fix' another error; SMEM is only for scalars/PRNG keys.","solutions":["Remove SMEM memory_space from tensor inputs; use default VMEM for arrays","Reserve SMEM for scalars and PRNG key data only"],"exampleFix":null,"handlingStrategy":"validation","validationCode":"# reject array-valued SMEM specs at setup time\nfor name, spec in specs.items():\n    if 'SMEM' in str(getattr(spec, 'memory_space', '') or '') and any(d != 1 for d in (spec.block_shape or ())):\n        raise ValueError(f'{name}: SMEM only supports scalar blocks')","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Reserve SMEM for scalars and PRNG keys only","Code-review memory_space choices per input"],"tags":["jax","pallas","tpu","smem","shape-mismatch"],"backgroundTag":"pallas-smem-scalar-only","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}