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, andPYTC_X_PANEL_BUDGET_GBpins 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]