{"record":{"id":"670d4eb842ab9543","repo":"jax-ml/jax","slug":"dma-source-destination-semaphore-arguments-must-be","errorCode":null,"errorMessage":"DMA source/destination/semaphore arguments must be Refs.","messagePattern":"DMA source/destination/semaphore arguments must be Refs\\.","errorType":"exception","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic/primitives.py","lineNumber":356,"sourceCode":"      priority=priority,\n      device_id_type=device_id_type,\n      add=add,\n  )\n  return []\ndma_start_p.to_lojax = _dma_start_to_lojax\n\n@dma_start_p.def_effectful_abstract_eval\ndef _dma_start_abstract_eval(*args, tree, device_id_type, priority, add):\n  if priority < 0:\n    raise ValueError(f\"DMA start priority must be non-negative: {priority}\")\n  src_ref_aval, dst_ref_aval, dst_sem_aval, src_sem_aval, device_id_aval = (\n      _dma_unflatten(tree, args)\n  )\n  if not all(\n      isinstance(x, (state.AbstractRef, state.TransformedRef))\n      for x in [src_ref_aval, dst_ref_aval, dst_sem_aval]\n  ):\n    raise ValueError(\n        \"DMA source/destination/semaphore arguments must be Refs.\")\n  dst_sem_shape = dst_sem_aval.shape\n  if dst_sem_shape:\n    raise ValueError(\n        f\"Cannot signal on a non-()-shaped semaphore: {dst_sem_shape}\"\n    )\n  if src_sem_aval is not None:\n    if not isinstance(src_sem_aval, (state.AbstractRef, state.TransformedRef)):\n      raise ValueError(\"DMA source semaphore must be a Ref.\")\n    src_sem_shape = src_sem_aval.shape\n    if src_sem_shape:\n      raise ValueError(\n          f\"Cannot signal on a non-()-shaped semaphore: {src_sem_shape}\"\n      )\n  return [], _get_dma_effects(\n      src_ref_aval,\n      dst_ref_aval,\n      dst_sem_aval,","sourceCodeStart":338,"sourceCodeEnd":374,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic/primitives.py#L338-L374","documentation":"dma_start operates on state.Ref arguments, not plain arrays, because the DMA engine writes memory in place and effects tracking depends on ref types. The abstract eval requires src_ref, dst_ref and dst_sem to be state.AbstractRef or state.TransformedRef; anything else (e.g. a jnp array or a triton-style pointer) is rejected.","triggerScenarios":"Calling dma_start with a plain jax.Array as source or destination (e.g. passing a block read out of a ref instead of the ref itself), or a non-ref semaphore.","commonSituations":"Confusing Pallas refs with arrays because they both support [] indexing — reading a block and passing it to dma_start; adapting example code that used different plumbing; forgetting to declare a semaphore via the ref API.","solutions":["Pass the Ref objects themselves, not blocks read from them","Ensure semaphores and buffers come from the kernel's ref parameters or state allocation","Check for stray get()/read calls between allocation and dma_start"],"exampleFix":"# before\nblock = src_ref[...]\ndma_start(block, dst_ref, sem_ref)  # block is an Array\n\n# after\ndma_start(src_ref, dst_ref, sem_ref)  # pass refs directly; index via the DMA's own offsets","handlingStrategy":"type-guard","validationCode":"from jax._src import state\nassert all(isinstance(v, (state.AbstractRef, state.TransformedRef)) or hasattr(v, 'shape') is False for v in (src_ref, dst_ref, dst_sem))","typeGuard":"def is_ref(v) -> bool:\n  from jax._src import state\n  return isinstance(getattr(v, 'aval', v), (state.AbstractRef, state.TransformedRef))","tryCatchPattern":null,"preventionTips":["Never index/read a ref before passing it to dma_start","Pass kernel ref parameters straight through to DMA ops"],"tags":["jax","pallas","mosaic","dma","ref-vs-array","type-validation"],"backgroundTag":"expected-ref-got-array","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}