pytc.solver.isdf_xtc_ccsd.contract_terms_t2_auto

pytc.solver.isdf_xtc_ccsd.contract_terms_t2_auto(t2, p, grad_p, u1, u3, d, x_backing, nocc, *, occupied_pair_batch_size=8, rank_panel_size=128, cap_bytes=25769803776)[source]

Three-tier X path, selected on measured free memory.

  • Tier 1 (fd_x_tier1_full_lift): X_vv fits the device-residency cap AND at most half the measured free device memory – the full-lift path (whole block device-lifted, one compiled rank scan), the fast path whenever the block fits.

  • Tier 2 (fd_x_tier2_host_resident): X_vv fails the device gate but fits in half the measured free host RAM – X_vv is lifted to host RAM once, then contracted by the pipelined panel loop with panel reads as host-array slices (the block fits a node’s RAM).

  • Tier 3 (fd_x_tier3_stream): otherwise – the same pipelined loop with panel reads from the backing (HDF5 dataset or ndarray).

Tiers 2/3 share the working-set panel loop: the panel width comes from the measured device working set and the next panel’s read + H2D transfer is prefetched behind the current kernel. The gate uses only measured free memory plus the cap; PYTC_X_FORCE_TIER = 1|2|3 (read at call time) pins the tier for tests and benchmarking, and PYTC_X_PANEL_BUDGET_GB pins the per-panel budget. Which tier fired is recorded in the tile_timers counters so receipts show it.

Parameters:
  • t2 (jax.Array)

  • p (jax.Array)

  • grad_p (jax.Array)

  • u1 (jax.Array)

  • u3 (jax.Array)

  • d (jax.Array)

  • nocc (int)

  • occupied_pair_batch_size (int)

  • rank_panel_size (int)

  • cap_bytes (int)

Return type:

Mapping[str, jax.Array]