[{"data":1,"prerenderedAt":1344},["ShallowReactive",2],{"navigation":3,"docyard:\u002Fapi\u002Ftransport\u002Fsamples\u002Fkernels":174,"api-content:\u002Fapi\u002Ftransport\u002Fsamples\u002Fkernels":1343,"docyard:crossref-index":1343},[4,8,101,165,170],{"title":5,"path":6,"stem":7},"Getting started","\u002Fgetting-started","1.getting-started",{"title":9,"path":10,"stem":11,"children":12},"Algorithms","\u002Falgorithms","2.algorithms",[13,15,58],{"title":9,"path":10,"stem":14},"2.algorithms\u002Findex",{"title":16,"path":17,"stem":18,"children":19},"Assignment","\u002Falgorithms\u002Fassignment","2.algorithms\u002F1.assignment\u002Findex",[20,21,25,29,32,36,40],{"title":16,"path":17,"stem":18},{"title":22,"path":23,"stem":24},"Quickstart","\u002Falgorithms\u002Fassignment\u002Fquickstart","2.algorithms\u002F1.assignment\u002F1.quickstart",{"title":26,"path":27,"stem":28},"Tracking","\u002Falgorithms\u002Fassignment\u002Ftracking","2.algorithms\u002F1.assignment\u002F2.tracking",{"title":9,"path":30,"stem":31},"\u002Falgorithms\u002Fassignment\u002Falgorithms","2.algorithms\u002F1.assignment\u002F3.algorithms",{"title":33,"path":34,"stem":35},"Reference","\u002Falgorithms\u002Fassignment\u002Freference","2.algorithms\u002F1.assignment\u002F4.reference",{"title":37,"path":38,"stem":39},"Choosing","\u002Falgorithms\u002Fassignment\u002Fchoosing","2.algorithms\u002F1.assignment\u002F5.choosing",{"title":41,"path":42,"stem":43,"children":44},"Tutorials","\u002Falgorithms\u002Fassignment\u002Ftutorials","2.algorithms\u002F1.assignment\u002F6.tutorials\u002Findex",[45,46,50,54],{"title":41,"path":42,"stem":43},{"title":47,"path":48,"stem":49},"Fundamentals","\u002Falgorithms\u002Fassignment\u002Ftutorials\u002Ffundamentals","2.algorithms\u002F1.assignment\u002F6.tutorials\u002F1.fundamentals",{"title":51,"path":52,"stem":53},"Backends","\u002Falgorithms\u002Fassignment\u002Ftutorials\u002Fbackends","2.algorithms\u002F1.assignment\u002F6.tutorials\u002F2.backends",{"title":55,"path":56,"stem":57},"Object tracking","\u002Falgorithms\u002Fassignment\u002Ftutorials\u002Ftracking","2.algorithms\u002F1.assignment\u002F6.tutorials\u002F3.tracking",{"title":59,"path":60,"stem":61,"children":62},"Transport","\u002Falgorithms\u002Ftransport","2.algorithms\u002F2.transport\u002Findex",[63,64,67,71,74,77,80,84],{"title":59,"path":60,"stem":61},{"title":22,"path":65,"stem":66},"\u002Falgorithms\u002Ftransport\u002Fquickstart","2.algorithms\u002F2.transport\u002F1.quickstart",{"title":68,"path":69,"stem":70},"Point-cloud tutorial","\u002Falgorithms\u002Ftransport\u002Fpoint-clouds","2.algorithms\u002F2.transport\u002F2.point-clouds",{"title":9,"path":72,"stem":73},"\u002Falgorithms\u002Ftransport\u002Falgorithms","2.algorithms\u002F2.transport\u002F3.algorithms",{"title":33,"path":75,"stem":76},"\u002Falgorithms\u002Ftransport\u002Freference","2.algorithms\u002F2.transport\u002F4.reference",{"title":37,"path":78,"stem":79},"\u002Falgorithms\u002Ftransport\u002Fchoosing","2.algorithms\u002F2.transport\u002F5.choosing",{"title":81,"path":82,"stem":83},"Building","\u002Falgorithms\u002Ftransport\u002Fbuilding","2.algorithms\u002F2.transport\u002F6.building",{"title":41,"path":85,"stem":86,"children":87},"\u002Falgorithms\u002Ftransport\u002Ftutorials","2.algorithms\u002F2.transport\u002F7.tutorials\u002Findex",[88,89,93,97],{"title":41,"path":85,"stem":86},{"title":90,"path":91,"stem":92},"Optimal transport","\u002Falgorithms\u002Ftransport\u002Ftutorials\u002Foptimal-transport","2.algorithms\u002F2.transport\u002F7.tutorials\u002F1.optimal-transport",{"title":94,"path":95,"stem":96},"Sinkhorn","\u002Falgorithms\u002Ftransport\u002Ftutorials\u002Fsinkhorn","2.algorithms\u002F2.transport\u002F7.tutorials\u002F2.sinkhorn",{"title":98,"path":99,"stem":100},"Point clouds","\u002Falgorithms\u002Ftransport\u002Ftutorials\u002Fpoint-clouds","2.algorithms\u002F2.transport\u002F7.tutorials\u002F3.point-clouds",{"title":102,"path":103,"stem":104,"children":105},"Resources","\u002Fresources","3.resources",[106,108,147,151,155],{"title":102,"path":103,"stem":107},"3.resources\u002Findex",{"title":41,"path":109,"stem":110,"children":111},"\u002Fresources\u002Ftutorials","3.resources\u002F1.tutorials\u002Findex",[112,113,131],{"title":41,"path":109,"stem":110},{"title":16,"path":114,"stem":115,"children":116,"page":130},"\u002Fresources\u002Ftutorials\u002Fassignment","3.resources\u002F1.tutorials\u002Fassignment",[117,122,126],{"title":118,"path":119,"stem":120,"icon":121},"Tutorial 1 — The Assignment Problem","\u002Fresources\u002Ftutorials\u002Fassignment\u002F01_the_assignment_problem","3.resources\u002F1.tutorials\u002Fassignment\u002F01_the_assignment_problem","i-lucide-notebook",{"title":123,"path":124,"stem":125,"icon":121},"Tutorial 2 — Backends and Batching","\u002Fresources\u002Ftutorials\u002Fassignment\u002F02_backends_and_batching","3.resources\u002F1.tutorials\u002Fassignment\u002F02_backends_and_batching",{"title":127,"path":128,"stem":129,"icon":121},"Tutorial 3 — Object Tracking with the Assignment Problem","\u002Fresources\u002Ftutorials\u002Fassignment\u002F03_object_tracking","3.resources\u002F1.tutorials\u002Fassignment\u002F03_object_tracking",false,{"title":59,"path":132,"stem":133,"children":134,"page":130},"\u002Fresources\u002Ftutorials\u002Ftransport","3.resources\u002F1.tutorials\u002Ftransport",[135,139,143],{"title":136,"path":137,"stem":138,"icon":121},"Tutorial 1 — What Is Optimal Transport?","\u002Fresources\u002Ftutorials\u002Ftransport\u002F01_optimal_transport","3.resources\u002F1.tutorials\u002Ftransport\u002F01_optimal_transport",{"title":140,"path":141,"stem":142,"icon":121},"Tutorial 2 — The Sinkhorn Algorithm","\u002Fresources\u002Ftutorials\u002Ftransport\u002F02_sinkhorn_algorithm","3.resources\u002F1.tutorials\u002Ftransport\u002F02_sinkhorn_algorithm",{"title":144,"path":145,"stem":146,"icon":121},"Tutorial 3 — Point-Cloud OT and Shape Learning","\u002Fresources\u002Ftutorials\u002Ftransport\u002F03_point_clouds","3.resources\u002F1.tutorials\u002Ftransport\u002F03_point_clouds",{"title":148,"path":149,"stem":150},"Assignment applications","\u002Fresources\u002Fassignment-applications","3.resources\u002F2.assignment-applications",{"title":152,"path":153,"stem":154},"Transport applications","\u002Fresources\u002Ftransport-applications","3.resources\u002F3.transport-applications",{"title":156,"path":157,"stem":158,"children":159},"Benchmarks","\u002Fresources\u002Fbenchmarks","3.resources\u002F4.benchmarks\u002Findex",[160,161],{"title":156,"path":157,"stem":158},{"title":162,"path":163,"stem":164},"Contributing benchmarks","\u002Fresources\u002Fbenchmarks\u002Fcontributing","3.resources\u002F4.benchmarks\u002Fcontributing",{"title":166,"path":167,"stem":168,"icon":169},"API Reference","\u002Fapi","4.api","i-lucide-package",{"title":171,"path":172,"stem":173},"References","\u002Freferences","5.references",{"docyard":175,"language":176,"name":177,"version":178,"items":179},"1","python","torchmatch","unknown",[180,188,197,208,217,226,233,241,250,258,267,276,286,295,305,313,322,331,339,346,353,360,367,390,397,405,412,419,426,433,442,505,513,522,531,538,545,552,558,565,572,602,638,677,686,709,716,722,728,734,741,748,755,762,769,775,782,789,796,802,809,815,822,833,842,851,860,869,878,887,896,905,914,923,930,935,939,944,949,953,999,1008,1015,1022,1029,1038,1047,1056,1065,1074,1083,1092,1101,1110,1118,1127,1136,1144,1152,1175,1199,1205,1210,1217,1223,1229,1235,1255,1261,1269,1275,1282,1289,1296,1303,1309,1316,1323,1330,1336],{"id":177,"kind":181,"path":182,"signature":183,"summary":184,"source":185},"module",[177],"module torchmatch","torchmatch - assignment and transport solvers for PyTorch.",{"file":186,"line":187},"torchmatch\u002F__init__.py",1,{"id":189,"kind":181,"path":190,"signature":192,"summary":193,"description":194,"source":195,"parent":177},"torchmatch.bench",[177,191],"bench","module torchmatch.bench","Benchmark collection CLI.","Helpers to capture a pytest-benchmark sweep on the local machine,\nscrub host-identifying fields, and write the result under\n``benchmarks\u002Fresults\u002F\u003Cslug>\u002F`` with a canonical filename that\nencodes the torchmatch \u002F Python \u002F PyTorch versions and the CUDA\nwheel variant.\n\nThe full contributor flow is :func:`init_machine` (once) +\n:func:`collect_run` (once per release). The aggregator under\n``scripts\u002Fbenchmark_aggregate.py`` then turns the per-machine\nfiles into the JSON datasets that the docs site consumes.",{"file":196,"line":187},"torchmatch\u002Fbench\u002F__init__.py",{"id":198,"kind":199,"path":200,"signature":202,"summary":203,"description":204,"source":205,"parent":189},"torchmatch.bench.collect_run","function",[177,191,201],"collect_run","def collect_run(results_root: Path, slug: str | None = None, tests_dir: Path | None = None, extra_pytest_args: Sequence[str] = (), run_id: str | None = None) -> Path","Run the benchmark sweep and write a scrubbed JSON to the slug dir.","``slug`` defaults to a unique directory under ``results_root``\ncontaining a ``machine.json`` (auto-detection). ``tests_dir`` defaults\nto ``\u003Ccwd>\u002Ftests``. Returns the path of the written file.",{"file":206,"line":207},"torchmatch\u002Fbench\u002F_collect.py",35,{"id":209,"kind":199,"path":210,"signature":212,"summary":213,"source":214,"parent":189},"torchmatch.bench.current_run_id",[177,191,211],"current_run_id","def current_run_id() -> str","Filename (sans .json) for a run captured right now on this interpreter.",{"file":215,"line":216},"torchmatch\u002Fbench\u002F_filename.py",73,{"id":218,"kind":199,"path":219,"signature":221,"summary":222,"description":223,"source":224,"parent":189},"torchmatch.bench.detect_variant",[177,191,220],"detect_variant","def detect_variant() -> str","Map the running PyTorch build to a wheel-variant tag.","Returns 'cpu' for CPU-only builds; 'cuXYZ' otherwise (matches the\nwheel suffixes produced by release.yml).",{"file":215,"line":225},22,{"id":227,"kind":199,"path":228,"signature":230,"source":231,"parent":189},"torchmatch.bench.make_filename",[177,191,229],"make_filename","def make_filename(version: str, py: str, torch_ver: str, variant: str, ts: str) -> str",{"file":215,"line":232},62,{"id":234,"kind":199,"path":235,"signature":237,"summary":238,"source":239,"parent":189},"torchmatch.bench.parse_filename",[177,191,236],"parse_filename","def parse_filename(name: str) -> dict | None","Reverse a canonical filename into its component dict, or None.",{"file":215,"line":240},85,{"id":242,"kind":199,"path":243,"signature":245,"summary":246,"source":247,"parent":189},"torchmatch.bench.build_slug",[177,191,244],"build_slug","def build_slug(machine_type: MachineType, cpu_short: str, gpu_short: str | None) -> str","Compose the flat slug used as the per-machine directory name.",{"file":248,"line":249},"torchmatch\u002Fbench\u002F_machine.py",140,{"id":251,"kind":199,"path":252,"signature":254,"summary":255,"source":256,"parent":189},"torchmatch.bench.detect_machine",[177,191,253],"detect_machine","def detect_machine(machine_type: MachineType) -> dict","Full machine.json payload for the local box.",{"file":248,"line":257},150,{"id":259,"kind":199,"path":260,"signature":262,"summary":263,"description":264,"source":265,"parent":189},"torchmatch.bench.init_machine",[177,191,261],"init_machine","def init_machine(results_root: Path, machine_type: MachineType | None = None, submitted_by: str = '', notes: str = '', interactive: bool = True, force: bool = False) -> Path","Write benchmarks\u002Fresults\u002F\u003Cslug>\u002Fmachine.json and return its path.","If a machine.json already exists at the derived slug, refuses to\noverwrite unless ``force=True``.",{"file":248,"line":266},189,{"id":268,"kind":199,"path":269,"signature":271,"summary":272,"description":273,"source":274,"parent":189},"torchmatch.bench.short_name",[177,191,270],"short_name","def short_name(brand: str) -> str","Normalize a noisy CPU\u002FGPU brand string to a filesystem-safe short slug.","Drops marketing tokens (\"(R)\", \"(TM)\", \"Core\", \"Processor\", \"Laptop GPU\",\n\"GeForce\"...) and runs of non-alphanumerics. Lowercased.\n\nExamples:\n    '12th Gen Intel(R) Core(TM) i7-12700H' -> 'intel-i7-12700h'\n    'AMD Ryzen 9 9950X 16-Core Processor'  -> 'amd-ryzen-9-9950x'\n    'NVIDIA RTX A1000 Laptop GPU'          -> 'nvidia-rtx-a1000'",{"file":248,"line":275},33,{"id":277,"kind":199,"path":278,"signature":280,"summary":281,"description":282,"source":283,"parent":189},"torchmatch.bench.scrub_machine_info",[177,191,279],"scrub_machine_info","def scrub_machine_info(data: dict, extra_tokens: list[str] | None = None) -> dict","Drop hostname\u002Fkernel fields in-place and redact stray hostname strings.","Returns the same dict for convenience. Idempotent.",{"file":284,"line":285},"torchmatch\u002Fbench\u002F_scrub.py",40,{"id":287,"kind":181,"path":288,"signature":290,"summary":291,"description":292,"source":293,"parent":177},"torchmatch.transport",[177,289],"transport","module torchmatch.transport","Continuous optimal-transport (OT) solvers.","Two sub-packages:\n\n- :mod:`torchmatch.transport.matrix`: cost-matrix in, plan \u002F divergence\n  out. The primary dispatcher; mirrors :mod:`torchmatch.assignment`.\n- :mod:`torchmatch.transport.samples`: point-clouds in, scalar loss\n  out. Backed by Triton kernels (CUDA-only).\n\nSee ``docs\u002Fsuperpowers\u002Fspecs\u002F2026-05-21-transport-ops-design.md``.",{"file":294,"line":187},"torchmatch\u002Ftransport\u002F__init__.py",{"id":296,"kind":199,"path":297,"signature":299,"summary":300,"description":301,"source":302,"parent":287},"torchmatch.transport.load_cpu",[177,289,298],"load_cpu","def load_cpu() -> None","Register ``torch.ops.transport.*`` (CPU backend).","Always registers the Python-side Sinkhorn custom_ops; additionally\nloads the C++ ``exact_emd`` extension when a prebuilt ``.so`` is\nshipped in the wheel or a C++ toolchain is available for the JIT\nfallback. With ``TORCHMATCH_SKIP_TRANSPORT=1`` at build time and no\ntoolchain at install, the extension is absent and EXACT_EMD fails\nwith a targeted RuntimeError at call time; the Sinkhorn backends\nkeep working. Set ``TORCHMATCH_DEBUG_LOADER=1`` to surface the\nunderlying JIT failure when diagnosing a partially-broken extension.",{"file":303,"line":304},"torchmatch\u002Ftransport\u002Fmatrix\u002F_cpu.py",55,{"id":306,"kind":199,"path":307,"signature":309,"source":310,"parent":287},"torchmatch.transport.load_cuda",[177,289,308],"load_cuda","def load_cuda() -> None",{"file":311,"line":312},"torchmatch\u002Ftransport\u002Fmatrix\u002F_cuda.py",24,{"id":314,"kind":181,"path":315,"signature":317,"summary":318,"description":319,"source":320,"parent":287},"torchmatch.transport.matrix",[177,289,316],"matrix","module torchmatch.transport.matrix","Cost-matrix optimal-transport solvers.","Takes a (B, N, M) cost matrix and returns a transport plan (or scalar\ndivergence). The dispatcher pattern mirrors :mod:`torchmatch.assignment`.\n\nSee ``docs\u002Fsuperpowers\u002Fspecs\u002F2026-05-21-transport-ops-design.md``.",{"file":321,"line":187},"torchmatch\u002Ftransport\u002Fmatrix\u002F__init__.py",{"id":323,"kind":324,"path":325,"signature":327,"source":328,"parent":314},"torchmatch.transport.matrix.Backend","type",[177,289,316,326],"Backend","class Backend(StrEnum):",{"file":329,"line":330},"torchmatch\u002Ftransport\u002Fmatrix\u002F_solve.py",43,{"id":332,"kind":333,"path":334,"signature":336,"source":337,"parent":323},"torchmatch.transport.matrix.Backend.AUTO","constant",[177,289,316,326,335],"AUTO","AUTO = 'auto'",{"file":329,"line":338},44,{"id":340,"kind":333,"path":341,"signature":343,"source":344,"parent":323},"torchmatch.transport.matrix.Backend.LOG_SINKHORN",[177,289,316,326,342],"LOG_SINKHORN","LOG_SINKHORN = 'log_sinkhorn'",{"file":329,"line":345},45,{"id":347,"kind":333,"path":348,"signature":350,"source":351,"parent":323},"torchmatch.transport.matrix.Backend.SINKHORN_DIVERGENCE",[177,289,316,326,349],"SINKHORN_DIVERGENCE","SINKHORN_DIVERGENCE = 'sinkhorn_divergence'",{"file":329,"line":352},46,{"id":354,"kind":333,"path":355,"signature":357,"source":358,"parent":323},"torchmatch.transport.matrix.Backend.UNBALANCED_SINKHORN",[177,289,316,326,356],"UNBALANCED_SINKHORN","UNBALANCED_SINKHORN = 'unbalanced_sinkhorn'",{"file":329,"line":359},47,{"id":361,"kind":333,"path":362,"signature":364,"source":365,"parent":323},"torchmatch.transport.matrix.Backend.EXACT_EMD",[177,289,316,326,363],"EXACT_EMD","EXACT_EMD = 'exact_emd'",{"file":329,"line":366},48,{"id":368,"kind":199,"path":369,"signature":371,"summary":372,"description":373,"params":374,"returns":385,"source":388,"parent":314},"torchmatch.transport.matrix.marginal_error",[177,289,316,370],"marginal_error","def marginal_error(log_plan: torch.Tensor, a: torch.Tensor, b: torch.Tensor) -> tuple[float, float]","Compute the max marginal error of a log-domain transport plan.","Useful for checking Sinkhorn convergence quality after ``solve()``.\nWorks with any backend that returns a log-plan (``LOG_SINKHORN``,\n``UNBALANCED_SINKHORN``, ``EXACT_EMD``).",[375,379,382],{"name":376,"type":377,"doc":378},"log_plan","torch.Tensor","Log-domain transport plan (B, N, M) or (N, M), as returned by\n``solve()``.",{"name":380,"type":377,"doc":381},"a","Source marginals (B, N) or (N,). Must match the ``a`` passed to\n``solve()``, or uniform weights if ``a=None`` was used.",{"name":383,"type":377,"doc":384},"b","Target marginals (B, M) or (M,).",{"type":386,"doc":387},"(row_err, col_err)","Max absolute deviation of plan row-sums and col-sums from ``a``\nand ``b``, respectively.",{"file":329,"line":389},231,{"id":391,"kind":199,"path":392,"signature":394,"source":395,"parent":314},"torchmatch.transport.matrix.solve",[177,289,316,393],"solve","def solve(cost: torch.Tensor, backend: Backend | str = Backend.AUTO, reg: float = 0.1, n_iter: int = 100, mask: torch.Tensor | None = None, a: torch.Tensor | None = None, b: torch.Tensor | None = None, scaling: float | None = None, rho: float = 1.0, cost_aa: torch.Tensor | None = None, cost_bb: torch.Tensor | None = None, unpack: bool = False) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor | None, torch.Tensor | None]",{"file":329,"line":396},98,{"id":398,"kind":181,"path":399,"signature":401,"summary":402,"source":403,"parent":314},"torchmatch.transport.matrix.ops",[177,289,316,400],"ops","module torchmatch.transport.matrix.ops","Direct handles for torch.ops.transport.* matrix-face ops.",{"file":404,"line":187},"torchmatch\u002Ftransport\u002Fmatrix\u002Fops.py",{"id":406,"kind":199,"path":407,"signature":409,"source":410,"parent":398},"torchmatch.transport.matrix.ops.exact_emd",[177,289,316,400,408],"exact_emd","def exact_emd(cost: torch.Tensor, mask: torch.Tensor | None = None, a: torch.Tensor | None = None, b: torch.Tensor | None = None) -> torch.Tensor",{"file":404,"line":411},3,{"id":413,"kind":199,"path":414,"signature":416,"source":417,"parent":398},"torchmatch.transport.matrix.ops.log_sinkhorn",[177,289,316,400,415],"log_sinkhorn","def log_sinkhorn(cost: torch.Tensor, eps: float, n_iter: int, a: torch.Tensor, b: torch.Tensor, mask: torch.Tensor | None = None, scaling: float | None = None) -> torch.Tensor",{"file":404,"line":418},9,{"id":420,"kind":199,"path":421,"signature":423,"source":424,"parent":398},"torchmatch.transport.matrix.ops.sinkhorn_divergence",[177,289,316,400,422],"sinkhorn_divergence","def sinkhorn_divergence(cost: torch.Tensor, eps: float, n_iter: int, a: torch.Tensor, b: torch.Tensor, mask: torch.Tensor | None = None, scaling: float | None = None, cost_aa: torch.Tensor | None = None, cost_bb: torch.Tensor | None = None) -> torch.Tensor",{"file":404,"line":425},18,{"id":427,"kind":199,"path":428,"signature":430,"source":431,"parent":398},"torchmatch.transport.matrix.ops.unbalanced_sinkhorn",[177,289,316,400,429],"unbalanced_sinkhorn","def unbalanced_sinkhorn(cost: torch.Tensor, eps: float, n_iter: int, rho: float, a: torch.Tensor, b: torch.Tensor, mask: torch.Tensor | None = None, scaling: float | None = None) -> torch.Tensor",{"file":404,"line":432},29,{"id":434,"kind":181,"path":435,"signature":437,"summary":438,"description":439,"source":440,"parent":287},"torchmatch.transport.samples",[177,289,436],"samples","module torchmatch.transport.samples","Samples-face optimal transport.","Point clouds in, scalar OT cost \u002F divergence out. Requires CUDA +\nTriton. See ``docs\u002Fsuperpowers\u002Fspecs\u002F2026-05-21-transport-ops-design.md``.",{"file":441,"line":187},"torchmatch\u002Ftransport\u002Fsamples\u002F__init__.py",{"id":443,"kind":199,"path":444,"signature":446,"summary":447,"params":448,"returns":500,"source":502,"parent":434},"torchmatch.transport.samples.loss",[177,289,436,445],"loss","def loss(x: torch.Tensor, y: torch.Tensor, blur: float = 0.05, debias: bool = False, reach: float | None = None, reach_x: float | None = None, reach_y: float | None = None, a: torch.Tensor | None = None, b: torch.Tensor | None = None, scaling: float = 0.5, n_iter: int | None = None, threshold: float | None = None, half_cost: bool = False, p: int = 2) -> torch.Tensor","Scalar OT cost \u002F divergence over point clouds (CUDA only).",[449,452,455,460,465,470,473,476,479,481,485,489,492,495],{"name":450,"type":377,"doc":451},"x","Source point cloud (n, d), float32, CUDA.",{"name":453,"type":377,"doc":454},"y","Target point cloud (m, d), float32, CUDA.",{"name":456,"type":457,"default":458,"doc":459},"blur","float","0.05","Bandwidth parameter; the entropy regularization is ``eps = blur**2``.",{"name":461,"type":462,"default":463,"doc":464},"debias","bool","False","If True, returns the Sinkhorn divergence\n``S_eps(x,y) - 0.5*S_eps(x,x) - 0.5*S_eps(y,y)`` which vanishes\nwhen ``x == y`` (unbiased). Requires three Sinkhorn solves.",{"name":466,"type":467,"default":468,"doc":469},"reach","float | None","None","Unbalanced OT: KL marginal penalty ``rho = reach**2`` applied to\nboth source and target. ``None`` = balanced OT.",{"name":471,"type":467,"default":468,"doc":472},"reach_x","Semi-unbalanced: KL penalty for the source marginal only.",{"name":474,"type":467,"default":468,"doc":475},"reach_y","Semi-unbalanced: KL penalty for the target marginal only.",{"name":380,"type":477,"default":468,"doc":478},"torch.Tensor | None","Source weights (n,). Uniform if None.",{"name":383,"type":477,"default":468,"doc":480},"Target weights (m,). Uniform if None.",{"name":482,"type":457,"default":483,"doc":484},"scaling","0.5","Geometric decay factor for the epsilon schedule, in ``(0, 1)``.\nSmaller values converge faster but may be less numerically stable.",{"name":486,"type":487,"default":468,"doc":488},"n_iter","int | None","Not supported; always raises. Control iterations via ``scaling``.",{"name":490,"type":467,"default":468,"doc":491},"threshold","Early-stopping threshold on potential change. Incompatible with\n``torch.compile`` (forces a host sync each check).",{"name":493,"type":462,"default":463,"doc":494},"half_cost","If True, uses ``cost = 0.5 * ||x - y||²`` instead of ``||x - y||²``.",{"name":496,"type":497,"default":498,"doc":499},"p","int","2","Cost exponent. Only ``p=2`` is supported.",{"type":445,"doc":501},"Scalar OT cost (shape ``()`` for 2-D input; ``(B,)`` for 3-D batch).",{"file":503,"line":504},"torchmatch\u002Ftransport\u002Fsamples\u002F_loss.py",52,{"id":506,"kind":181,"path":507,"signature":509,"summary":510,"source":511,"parent":434},"torchmatch.transport.samples.kernels",[177,289,436,508],"kernels","module torchmatch.transport.samples.kernels","Triton kernels for streaming optimal transport.",{"file":512,"line":187},"torchmatch\u002Ftransport\u002Fsamples\u002Fkernels\u002F__init__.py",{"id":514,"kind":181,"path":515,"signature":517,"summary":518,"description":519,"source":520,"parent":506},"torchmatch.transport.samples.kernels.apply_sqeuclid",[177,289,436,508,516],"apply_sqeuclid","module torchmatch.transport.samples.kernels.apply_sqeuclid","Backward-compatible re-exports for apply kernels.","The actual implementations live in:\n- apply_raw.py: raw-form kernels (mat5, deprecated vec\u002Fmat wrappers)\n- apply_shifted.py: shifted-form kernels (vec\u002Fmat with shifted potentials)",{"file":521,"line":187},"torchmatch\u002Ftransport\u002Fsamples\u002Fkernels\u002Fapply_sqeuclid.py",{"id":523,"kind":181,"path":524,"signature":526,"summary":527,"description":528,"source":529,"parent":506},"torchmatch.transport.samples.kernels.cg_dense",[177,289,436,508,525],"cg_dense","module torchmatch.transport.samples.kernels.cg_dense","Dense CG solver with cached transport plan matrix.","This module provides a CG solver that materializes the O(n*m) transport plan\nmatrix P and uses dense matrix-vector products (P @ v, P.T @ v) instead of\nstreaming Triton kernels. This is significantly faster for small problem sizes\n(n \u003C= 8192) where:\n1. The O(n²) plan fits in GPU memory\n2. Kernel launch overhead dominates streaming computation\n3. Dense matvecs can leverage tensor core saturation\n\nPerformance comparison (n=2048, d=64):\n- Streaming CG (Triton): ~48 ms (64 kernel launches × 0.76 ms)\n- Dense CG (cached P):   ~5.9 ms (8.12x speedup!)\n\nMemory trade-off:\n- Streaming: O(nd) memory\n- Dense: O(nm) memory for transport plan P\n\nCrossover point: n ≈ 8192-10000 where memory becomes the bottleneck.",{"file":530,"line":187},"torchmatch\u002Ftransport\u002Fsamples\u002Fkernels\u002Fcg_dense.py",{"id":532,"kind":324,"path":533,"signature":535,"summary":536,"source":537,"parent":523},"torchmatch.transport.samples.kernels.cg_dense.DenseCgInfo",[177,289,436,508,525,534],"DenseCgInfo","class DenseCgInfo:","Convergence information for dense CG solver.",{"file":530,"line":432},{"id":539,"kind":540,"path":541,"signature":543,"source":544,"parent":532},"torchmatch.transport.samples.kernels.cg_dense.DenseCgInfo.cg_converged","property",[177,289,436,508,525,534,542],"cg_converged","cg_converged: bool",{"file":530,"line":275},{"id":546,"kind":540,"path":547,"signature":549,"source":550,"parent":532},"torchmatch.transport.samples.kernels.cg_dense.DenseCgInfo.cg_iters",[177,289,436,508,525,534,548],"cg_iters","cg_iters: int",{"file":530,"line":551},34,{"id":553,"kind":540,"path":554,"signature":556,"source":557,"parent":532},"torchmatch.transport.samples.kernels.cg_dense.DenseCgInfo.cg_residual",[177,289,436,508,525,534,555],"cg_residual","cg_residual: float",{"file":530,"line":207},{"id":559,"kind":540,"path":560,"signature":562,"source":563,"parent":532},"torchmatch.transport.samples.kernels.cg_dense.DenseCgInfo.cg_initial_residual",[177,289,436,508,525,534,561],"cg_initial_residual","cg_initial_residual: float",{"file":530,"line":564},36,{"id":566,"kind":567,"path":568,"signature":570,"source":571,"parent":532},"torchmatch.transport.samples.kernels.cg_dense.DenseCgInfo.__init__","method",[177,289,436,508,525,534,569],"__init__","def __init__(self, cg_converged: bool, cg_iters: int, cg_residual: float, cg_initial_residual: float) -> None",{"file":530,"line":187},{"id":573,"kind":199,"path":574,"signature":576,"summary":577,"params":578,"returns":598,"source":600,"parent":523},"torchmatch.transport.samples.kernels.cg_dense.materialize_transport_plan",[177,289,436,508,525,575],"materialize_transport_plan","def materialize_transport_plan(x: torch.Tensor, y: torch.Tensor, f_hat: torch.Tensor, g_hat: torch.Tensor, eps: float, x2: torch.Tensor | None = None, y2: torch.Tensor | None = None) -> torch.Tensor","Materialize the transport plan P = exp((f_hat + g_hat - C) \u002F eps).",[579,581,583,586,589,592,595],{"name":450,"type":377,"doc":580},"Source points (n, d)",{"name":453,"type":377,"doc":582},"Target points (m, d)",{"name":584,"type":377,"doc":585},"f_hat","Source raw-form potential (n,)",{"name":587,"type":377,"doc":588},"g_hat","Target raw-form potential (m,)",{"name":590,"type":457,"doc":591},"eps","Entropy regularization",{"name":593,"type":477,"default":468,"doc":594},"x2","Precomputed ||x_i||² (n,), optional",{"name":596,"type":477,"default":468,"doc":597},"y2","Precomputed ||y_j||² (m,), optional",{"type":377,"doc":599},"Transport plan matrix (n, m), dtype=float32",{"file":530,"line":601},39,{"id":603,"kind":199,"path":604,"signature":606,"summary":607,"description":608,"params":609,"returns":633,"source":636,"parent":523},"torchmatch.transport.samples.kernels.cg_dense.dense_cg_solve",[177,289,436,508,525,605],"dense_cg_solve","def dense_cg_solve(P: torch.Tensor, diag_x: torch.Tensor, denom: torch.Tensor, rhs: torch.Tensor, max_iter: int = 100, rtol: float = 1e-06, atol: float = 1e-06) -> tuple[torch.Tensor, DenseCgInfo]","Solve the Schur complement system using dense matrix operations.","Solves: (denom * I - P.T @ diag(1\u002Fdiag_x) @ P) @ z = rhs\n\nThis is the inner CG solve for the HVP. Instead of using streaming Triton\nkernels to apply the transport plan, we cache P and use dense matvecs.",[610,613,616,619,622,626,630],{"name":611,"type":377,"doc":612},"P","Materialized transport plan (n, m).",{"name":614,"type":377,"doc":615},"diag_x","Diagonal scaling for source marginal (n,).",{"name":617,"type":377,"doc":618},"denom","Diagonal of the Schur complement (m,), includes regularization.",{"name":620,"type":377,"doc":621},"rhs","Right-hand side vector (m,).",{"name":623,"type":497,"default":624,"doc":625},"max_iter","100","Maximum CG iterations.",{"name":627,"type":457,"default":628,"doc":629},"rtol","1e-06","Relative tolerance.",{"name":631,"type":457,"default":628,"doc":632},"atol","Absolute tolerance.",{"type":634,"doc":635},"z","Solution vector (m,).",{"file":530,"line":637},91,{"id":639,"kind":199,"path":640,"signature":642,"summary":643,"description":644,"params":645,"returns":673,"source":675,"parent":523},"torchmatch.transport.samples.kernels.cg_dense.hvp_dense_cg",[177,289,436,508,525,641],"hvp_dense_cg","def hvp_dense_cg(x: torch.Tensor, y: torch.Tensor, f_hat: torch.Tensor, g_hat: torch.Tensor, A: torch.Tensor, eps: float, rho_x: float | None = None, rho_y: float | None = None, tau2: float = 1e-05, max_cg_iter: int = 100, cg_rtol: float = 1e-06, cg_atol: float = 1e-06) -> tuple[torch.Tensor, DenseCgInfo]","Compute HVP using dense CG with cached transport plan.","This is a drop-in replacement for the streaming Triton HVP that is\n8x faster for small problem sizes (n \u003C= 8192).",[646,647,648,649,650,653,654,657,660,664,667,670],{"name":450,"type":377,"doc":580},{"name":453,"type":377,"doc":582},{"name":584,"type":377,"doc":585},{"name":587,"type":377,"doc":588},{"name":651,"type":377,"doc":652},"A","Input matrix for HVP (n, d)",{"name":590,"type":457,"doc":591},{"name":655,"type":467,"default":468,"doc":656},"rho_x","Source marginal KL penalty (None = strict constraint)",{"name":658,"type":467,"default":468,"doc":659},"rho_y","Target marginal KL penalty (None = strict constraint)",{"name":661,"type":457,"default":662,"doc":663},"tau2","1e-05","Tikhonov regularization (only used for balanced OT)",{"name":665,"type":497,"default":624,"doc":666},"max_cg_iter","Maximum CG iterations",{"name":668,"type":457,"default":628,"doc":669},"cg_rtol","Relative CG tolerance",{"name":671,"type":457,"default":628,"doc":672},"cg_atol","Absolute CG tolerance",{"type":377,"doc":674},"Hessian-vector product H @ A, shape (n, d)",{"file":530,"line":676},207,{"id":678,"kind":324,"path":679,"signature":681,"summary":682,"description":683,"source":684,"parent":523},"torchmatch.transport.samples.kernels.cg_dense.CachedDenseHVP",[177,289,436,508,525,680],"CachedDenseHVP","class CachedDenseHVP:","Cached HVP context that materializes transport plan once.","This class provides significant speedup when computing multiple HVPs\nwith the same potentials (e.g., during outer Newton CG iterations).\nThe transport plan P is materialized once during __init__, and reused\nfor all subsequent hvp() calls.\n\nPerformance improvement: ~19x speedup for 19 CG iterations\n(avoids re-materializing O(n²) transport plan each iteration).",{"file":530,"line":685},387,{"id":687,"kind":567,"path":688,"signature":689,"summary":690,"params":691,"source":707,"parent":678},"torchmatch.transport.samples.kernels.cg_dense.CachedDenseHVP.__init__",[177,289,436,508,525,680,569],"def __init__(self, x: torch.Tensor, y: torch.Tensor, f_hat: torch.Tensor, g_hat: torch.Tensor, eps: float, rho_x: float | None = None, rho_y: float | None = None, tau2: float = 1e-05, max_cg_iter: int = 100, cg_rtol: float = 1e-06, cg_atol: float = 1e-06)","Initialize cached HVP context.",[692,693,694,695,696,697,699,701,703,705,706],{"name":450,"type":377,"doc":580},{"name":453,"type":377,"doc":582},{"name":584,"type":377,"doc":585},{"name":587,"type":377,"doc":588},{"name":590,"type":457,"doc":591},{"name":655,"type":467,"default":468,"doc":698},"Source marginal KL penalty (None = balanced)",{"name":658,"type":467,"default":468,"doc":700},"Target marginal KL penalty (None = balanced)",{"name":661,"type":457,"default":662,"doc":702},"Tikhonov regularization",{"name":665,"type":497,"default":624,"doc":704},"Maximum inner CG iterations",{"name":668,"type":457,"default":628,"doc":669},{"name":671,"type":457,"default":628,"doc":672},{"file":530,"line":708},405,{"id":710,"kind":540,"path":711,"signature":713,"source":714,"parent":678},"torchmatch.transport.samples.kernels.cg_dense.CachedDenseHVP.eps_f",[177,289,436,508,525,680,712],"eps_f","eps_f = float(eps)",{"file":530,"line":715},447,{"id":717,"kind":540,"path":718,"signature":719,"source":720,"parent":678},"torchmatch.transport.samples.kernels.cg_dense.CachedDenseHVP.max_cg_iter",[177,289,436,508,525,680,665],"max_cg_iter = max_cg_iter",{"file":530,"line":721},448,{"id":723,"kind":540,"path":724,"signature":725,"source":726,"parent":678},"torchmatch.transport.samples.kernels.cg_dense.CachedDenseHVP.cg_rtol",[177,289,436,508,525,680,668],"cg_rtol = cg_rtol",{"file":530,"line":727},449,{"id":729,"kind":540,"path":730,"signature":731,"source":732,"parent":678},"torchmatch.transport.samples.kernels.cg_dense.CachedDenseHVP.cg_atol",[177,289,436,508,525,680,671],"cg_atol = cg_atol",{"file":530,"line":733},450,{"id":735,"kind":540,"path":736,"signature":738,"source":739,"parent":678},"torchmatch.transport.samples.kernels.cg_dense.CachedDenseHVP.n",[177,289,436,508,525,680,737],"n","n = n",{"file":530,"line":740},454,{"id":742,"kind":540,"path":743,"signature":745,"source":746,"parent":678},"torchmatch.transport.samples.kernels.cg_dense.CachedDenseHVP.m",[177,289,436,508,525,680,744],"m","m = m",{"file":530,"line":747},455,{"id":749,"kind":540,"path":750,"signature":752,"source":753,"parent":678},"torchmatch.transport.samples.kernels.cg_dense.CachedDenseHVP.d",[177,289,436,508,525,680,751],"d","d = d",{"file":530,"line":754},456,{"id":756,"kind":540,"path":757,"signature":759,"source":760,"parent":678},"torchmatch.transport.samples.kernels.cg_dense.CachedDenseHVP.x_f",[177,289,436,508,525,680,758],"x_f","x_f = x.float()",{"file":530,"line":761},459,{"id":763,"kind":540,"path":764,"signature":766,"source":767,"parent":678},"torchmatch.transport.samples.kernels.cg_dense.CachedDenseHVP.y_f",[177,289,436,508,525,680,765],"y_f","y_f = y.float()",{"file":530,"line":768},460,{"id":770,"kind":333,"path":771,"signature":772,"source":773,"parent":678},"torchmatch.transport.samples.kernels.cg_dense.CachedDenseHVP.P",[177,289,436,508,525,680,611],"P = materialize_transport_plan(x, y, f_hat, g_hat, self.eps_f, x2, y2)",{"file":530,"line":774},467,{"id":776,"kind":540,"path":777,"signature":779,"source":780,"parent":678},"torchmatch.transport.samples.kernels.cg_dense.CachedDenseHVP.a_hat",[177,289,436,508,525,680,778],"a_hat","a_hat = torch.mv(self.P, ones_m)",{"file":530,"line":781},472,{"id":783,"kind":540,"path":784,"signature":786,"source":787,"parent":678},"torchmatch.transport.samples.kernels.cg_dense.CachedDenseHVP.b_hat",[177,289,436,508,525,680,785],"b_hat","b_hat = torch.mv(self.P.t(), ones_n)",{"file":530,"line":788},473,{"id":790,"kind":540,"path":791,"signature":793,"source":794,"parent":678},"torchmatch.transport.samples.kernels.cg_dense.CachedDenseHVP.is_balanced",[177,289,436,508,525,680,792],"is_balanced","is_balanced = rho_x is None and rho_y is None",{"file":530,"line":795},479,{"id":797,"kind":540,"path":798,"signature":799,"source":800,"parent":678},"torchmatch.transport.samples.kernels.cg_dense.CachedDenseHVP.diag_x",[177,289,436,508,525,680,614],"diag_x = diag_factor_x * a_hat_clamped",{"file":530,"line":801},496,{"id":803,"kind":540,"path":804,"signature":806,"source":807,"parent":678},"torchmatch.transport.samples.kernels.cg_dense.CachedDenseHVP.diag_y",[177,289,436,508,525,680,805],"diag_y","diag_y = diag_factor_y * b_hat_clamped",{"file":530,"line":808},497,{"id":810,"kind":540,"path":811,"signature":812,"source":813,"parent":678},"torchmatch.transport.samples.kernels.cg_dense.CachedDenseHVP.denom",[177,289,436,508,525,680,617],"denom = self.diag_y + self.eps_f * float(tau2)",{"file":530,"line":814},501,{"id":816,"kind":540,"path":817,"signature":819,"source":820,"parent":678},"torchmatch.transport.samples.kernels.cg_dense.CachedDenseHVP.Py",[177,289,436,508,525,680,818],"Py","Py = torch.mm(self.P, self.y_f)",{"file":530,"line":821},506,{"id":823,"kind":567,"path":824,"signature":826,"summary":827,"params":828,"returns":830,"source":831,"parent":678},"torchmatch.transport.samples.kernels.cg_dense.CachedDenseHVP.hvp",[177,289,436,508,525,680,825],"hvp","def hvp(self, A: torch.Tensor) -> tuple[torch.Tensor, DenseCgInfo]","Compute HVP using cached transport plan.",[829],{"name":651,"type":377,"doc":652},{"type":377,"doc":674},{"file":530,"line":832},508,{"id":834,"kind":181,"path":835,"signature":837,"summary":838,"description":839,"source":840,"parent":506},"torchmatch.transport.samples.kernels.streaming_sqeuclid",[177,289,436,508,836],"streaming_sqeuclid","module torchmatch.transport.samples.kernels.streaming_sqeuclid","Sinkhorn OT with streaming softmax and shifted potentials.","This module implements a reformulated Sinkhorn algorithm that aligns exactly with\nstreaming-softmax's interface, enabling potential future integration with optimized\nstreaming-softmax kernels.",{"file":841,"line":187},"torchmatch\u002Ftransport\u002Fsamples\u002Fkernels\u002Fstreaming_sqeuclid.py",{"id":843,"kind":199,"path":844,"signature":846,"summary":847,"description":848,"source":849,"parent":834},"torchmatch.transport.samples.kernels.streaming_sqeuclid.precompute_sinkhorn_inputs",[177,289,436,508,836,845],"precompute_sinkhorn_inputs","def precompute_sinkhorn_inputs(x: torch.Tensor, y: torch.Tensor, a: torch.Tensor, b: torch.Tensor, eps: float, cost_scale: float = 1.0) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]","Precompute static bias components.","Args:\n    x: Source points [n, d]\n    y: Target points [m, d]\n    a: Source marginal weights [n]\n    b: Target marginal weights [m]\n    eps: Regularization parameter\n    cost_scale: Cost scaling (1.0 for full ||x-y||², 0.5 for half ||x-y||²\u002F2)\n\nReturns:\n    alpha: Source squared norms [n] = cost_scale * ||x||²\n    beta: Target squared norms [m] = cost_scale * ||y||²\n    gamma: Scaled log source weights [n] = eps * log(a) (for g-update)\n    delta: Scaled log target weights [m] = eps * log(b) (for f-update)\n\nNotes:\n    - For cost_scale=1.0: full squared Euclidean C = ||x-y||²\n    - For cost_scale=0.5: half squared Euclidean C = ||x-y||²\u002F2\n    - Use sinkhorn_lse() with raw x, y coordinates (not pre-scaled Q, K)",{"file":841,"line":850},69,{"id":852,"kind":199,"path":853,"signature":855,"summary":856,"description":857,"source":858,"parent":834},"torchmatch.transport.samples.kernels.streaming_sqeuclid.compute_bias_f",[177,289,436,508,836,854],"compute_bias_f","def compute_bias_f(g: torch.Tensor, beta: torch.Tensor, delta: torch.Tensor, eps: float) -> torch.Tensor","Compute pre-scaled bias for f-update: u = (ĝ + δ)\u002Fε.","Args:\n    g: Current g potential [m]\n    beta: Target squared norms [m] = cost_scale * ||y||²\n    delta: Scaled log target weights [m] = eps * log(b)\n    eps: Regularization parameter\n\nReturns:\n    u: Pre-scaled bias [m] = (g - beta + delta) \u002F eps",{"file":841,"line":859},109,{"id":861,"kind":199,"path":862,"signature":864,"summary":865,"description":866,"source":867,"parent":834},"torchmatch.transport.samples.kernels.streaming_sqeuclid.compute_bias_g",[177,289,436,508,836,863],"compute_bias_g","def compute_bias_g(f: torch.Tensor, alpha: torch.Tensor, gamma: torch.Tensor, eps: float) -> torch.Tensor","Compute pre-scaled bias for g-update: v = (f̂ + γ)\u002Fε.","Args:\n    f: Current f potential [n]\n    alpha: Source squared norms [n] = cost_scale * ||x||²\n    gamma: Scaled log source weights [n] = eps * log(a)\n    eps: Regularization parameter\n\nReturns:\n    v: Pre-scaled bias [n] = (f - alpha + gamma) \u002F eps",{"file":841,"line":868},130,{"id":870,"kind":199,"path":871,"signature":873,"summary":874,"description":875,"source":876,"parent":834},"torchmatch.transport.samples.kernels.streaming_sqeuclid.sinkhorn_lse",[177,289,436,508,836,872],"sinkhorn_lse","def sinkhorn_lse(x: torch.Tensor, y: torch.Tensor, bias: torch.Tensor, eps: float, cost_scale: float = 1.0, damping: float = 1.0, allow_tf32: bool = True, use_exp2: bool = True, block_m: int | None = None, block_n: int | None = None, block_k: int | None = None, num_warps: int | None = None, num_stages: int = 2, autotune: bool = True) -> torch.Tensor","Compute shifted potential using the streaming-softmax kernel.","This computes: out_i = -ε * damping * LSE_j[coord_scale * x_i·y_j \u002F ε + bias_j]\n\nThe kernel applies coord_scale = 2 * cost_scale inside the kernel to scale\nthe dot product, ensuring consistent TF32 rounding between fused and separate\nkernel paths.\n\nArgs:\n    x: Source coordinates [n, d]\n    y: Target coordinates [m, d]\n    bias: Pre-scaled bias [m]\n    eps: Regularization parameter\n    cost_scale: Cost scaling (1.0 for full ||x-y||², 0.5 for half ||x-y||²\u002F2)\n    damping: Unbalanced OT damping (1.0 for balanced)\n    allow_tf32: Enable TF32 for matmul\n    use_exp2: Use exp2\u002Flog2 for better numerical stability\n    block_m, block_n, block_k: Manual block sizes (disables autotune)\n    num_warps: Number of warps (disables autotune)\n    num_stages: Number of pipeline stages\n    autotune: Enable autotuning\n\nReturns:\n    out: Shifted potential [n]",{"file":841,"line":877},945,{"id":879,"kind":199,"path":880,"signature":882,"summary":883,"description":884,"source":885,"parent":834},"torchmatch.transport.samples.kernels.streaming_sqeuclid.sinkhorn_lse_fused",[177,289,436,508,836,881],"sinkhorn_lse_fused","def sinkhorn_lse_fused(x: torch.Tensor, y: torch.Tensor, g_hat: torch.Tensor, log_w: torch.Tensor, eps: float, cost_scale: float = 1.0, damping: float = 1.0, allow_tf32: bool = True, use_exp2: bool = True, block_m: int | None = None, block_n: int | None = None, block_k: int | None = None, num_warps: int | None = None, num_stages: int = 2, autotune: bool = True) -> torch.Tensor","Fused LSE kernel that computes bias in SRAM (matches symmetric kernel interface).","This computes: out_i = -ε * damping * LSE_j[coord_scale * x_i·y_j \u002F ε + ĝ_j\u002Fε + log(w_j)]\n\nKEY OPTIMIZATION: Load ĝ and log(w) separately and compute bias = ĝ\u002Fε + log(w) in SRAM.\nThis matches the symmetric kernel interface and eliminates Python kernel launch overhead.\n\nArgs:\n    x: Source coordinates [n, d]\n    y: Target coordinates [m, d]\n    g_hat: Shifted potential [m] (ĝ = g - β for f-update)\n    log_w: Log marginal [m] (log(b) for f-update, NOT scaled by eps!)\n    eps: Regularization parameter\n    cost_scale: Cost scaling (1.0 for full, 0.5 for half)\n    damping: Unbalanced OT damping (1.0 for balanced)\n    allow_tf32: Enable TF32 for matmul\n    use_exp2: Use exp2\u002Flog2 for numerical stability\n    block_m, block_n, block_k, num_warps: Manual block sizes\n    num_stages: Pipeline stages\n    autotune: Enable autotuning\n\nReturns:\n    out: Shifted potential [n]",{"file":841,"line":886},1101,{"id":888,"kind":199,"path":889,"signature":891,"summary":892,"description":893,"source":894,"parent":834},"torchmatch.transport.samples.kernels.streaming_sqeuclid.sinkhorn_symmetric_step",[177,289,436,508,836,890],"sinkhorn_symmetric_step","def sinkhorn_symmetric_step(x: torch.Tensor, y: torch.Tensor, f_hat: torch.Tensor, g_hat: torch.Tensor, log_a: torch.Tensor, log_b: torch.Tensor, eps: float, cost_scale: float = 1.0, alpha: float = 0.5, damping_f: float = 1.0, damping_g: float = 1.0, allow_tf32: bool = True, use_exp2: bool = True, autotune: bool = True, block_m: int | None = None, block_n: int | None = None, block_k: int | None = None, num_warps: int | None = None, num_stages: int = 2, label_x: torch.Tensor | None = None, label_y: torch.Tensor | None = None, label_cost_matrix: torch.Tensor | None = None, lambda_x: float = 1.0, lambda_y: float = 0.0) -> tuple[torch.Tensor, torch.Tensor]","Fused symmetric Sinkhorn step: computes both f and g updates in ONE kernel.","This is the key optimization over separate f\u002Fg kernel calls - reduces kernel\nlaunch overhead by 50% and improves GPU occupancy.\n\nKEY: Uses x, y directly (no Q, K pre-allocation). The coord_scale = 2*cost_scale\nis applied inside the kernel, avoiding memory allocation overhead.\n\nArgs:\n    x: Source coordinates [n, d] (NOT pre-scaled!)\n    y: Target coordinates [m, d] (NOT pre-scaled!)\n    f_hat: Current shifted f potential [n]\n    g_hat: Current shifted g potential [m]\n    log_a: Log source weights [n]\n    log_b: Log target weights [m]\n    eps: Regularization parameter\n    cost_scale: Cost scaling (1.0 for full, 0.5 for half cost)\n    alpha: Averaging weight (0.5 for symmetric, 1.0 for full update)\n    damping_f: Unbalanced OT damping for f (1.0 for balanced)\n    damping_g: Unbalanced OT damping for g (1.0 for balanced)\n    allow_tf32: Enable TF32 for matmul\n    use_exp2: Use exp2\u002Flog2 optimization\n    autotune: Enable Triton autotuning (recommended)\n    block_m, block_n, block_k: Manual block sizes (disables autotune)\n    num_warps: Number of warps (disables autotune)\n    num_stages: Pipeline stages\n    label_x: int32\u002Fint64 labels for x [n] (OTDD)\n    label_y: int32\u002Fint64 labels for y [m] (OTDD)\n    label_cost_matrix: W [V, V] label distance matrix (OTDD)\n    lambda_x: Weight for Euclidean cost (default 1.0)\n    lambda_y: Weight for label cost (default 0.0 = no label cost)\n\nReturns:\n    f_hat_new, g_hat_new: Updated shifted potentials",{"file":841,"line":895},1251,{"id":897,"kind":199,"path":898,"signature":900,"summary":901,"description":902,"source":903,"parent":834},"torchmatch.transport.samples.kernels.streaming_sqeuclid.shifted_to_standard_potentials",[177,289,436,508,836,899],"shifted_to_standard_potentials","def shifted_to_standard_potentials(f_hat: torch.Tensor, g_hat: torch.Tensor, alpha: torch.Tensor, beta: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]","Convert shifted potentials back to standard form.","Args:\n    f_hat: Shifted source potential [n]\n    g_hat: Shifted target potential [m]\n    alpha: Source squared norms [n] = cost_scale * ||x||²\n    beta: Target squared norms [m] = cost_scale * ||y||²\n\nReturns:\n    f: Standard source potential [n] = f_hat + alpha\n    g: Standard target potential [m] = g_hat + beta",{"file":841,"line":904},1511,{"id":906,"kind":199,"path":907,"signature":909,"summary":910,"description":911,"source":912,"parent":834},"torchmatch.transport.samples.kernels.streaming_sqeuclid.standard_to_shifted_potentials",[177,289,436,508,836,908],"standard_to_shifted_potentials","def standard_to_shifted_potentials(f: torch.Tensor, g: torch.Tensor, alpha: torch.Tensor, beta: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]","Convert standard potentials to shifted form.","Args:\n    f: Standard source potential [n]\n    g: Standard target potential [m]\n    alpha: Source squared norms [n] = cost_scale * ||x||²\n    beta: Target squared norms [m] = cost_scale * ||y||²\n\nReturns:\n    f_hat: Shifted source potential [n] = f - alpha\n    g_hat: Shifted target potential [m] = g - beta",{"file":841,"line":913},1532,{"id":915,"kind":181,"path":916,"signature":918,"summary":919,"description":920,"source":921,"parent":506},"torchmatch.transport.samples.kernels.cg_python_batched",[177,289,436,508,917],"cg_python_batched","module torchmatch.transport.samples.kernels.cg_python_batched","Python-level batched CG for HVP acceleration.","This module implements a batched CG solver that uses the\napply_plan_vec_shifted kernels for improved performance and reduced memory\nusage.\n\nThis avoids the mysterious Triton compilation state issues observed with\nthe fully-inline batched CG kernel while still providing a clean interface\nfor batched CG solving.\n\nIncludes a torch.compile-compatible version for additional speedup.",{"file":922,"line":187},"torchmatch\u002Ftransport\u002Fsamples\u002Fkernels\u002Fcg_python_batched.py",{"id":924,"kind":324,"path":925,"signature":927,"summary":928,"source":929,"parent":915},"torchmatch.transport.samples.kernels.cg_python_batched.PythonBatchedCgInfo",[177,289,436,508,917,926],"PythonBatchedCgInfo","class PythonBatchedCgInfo:","Information about Python batched CG execution.",{"file":922,"line":312},{"id":931,"kind":540,"path":932,"signature":543,"source":933,"parent":924},"torchmatch.transport.samples.kernels.cg_python_batched.PythonBatchedCgInfo.cg_converged",[177,289,436,508,917,926,542],{"file":922,"line":934},28,{"id":936,"kind":540,"path":937,"signature":549,"source":938,"parent":924},"torchmatch.transport.samples.kernels.cg_python_batched.PythonBatchedCgInfo.cg_iters",[177,289,436,508,917,926,548],{"file":922,"line":432},{"id":940,"kind":540,"path":941,"signature":556,"source":942,"parent":924},"torchmatch.transport.samples.kernels.cg_python_batched.PythonBatchedCgInfo.cg_residual",[177,289,436,508,917,926,555],{"file":922,"line":943},30,{"id":945,"kind":540,"path":946,"signature":562,"source":947,"parent":924},"torchmatch.transport.samples.kernels.cg_python_batched.PythonBatchedCgInfo.cg_initial_residual",[177,289,436,508,917,926,561],{"file":922,"line":948},31,{"id":950,"kind":567,"path":951,"signature":570,"source":952,"parent":924},"torchmatch.transport.samples.kernels.cg_python_batched.PythonBatchedCgInfo.__init__",[177,289,436,508,917,926,569],{"file":922,"line":187},{"id":954,"kind":199,"path":955,"signature":957,"summary":958,"description":959,"params":960,"returns":995,"source":997,"parent":915},"torchmatch.transport.samples.kernels.cg_python_batched.python_batched_cg_solve",[177,289,436,508,917,956],"python_batched_cg_solve","def python_batched_cg_solve(x: torch.Tensor, y: torch.Tensor, f: torch.Tensor, g: torch.Tensor, diag_x: torch.Tensor, denom: torch.Tensor, rhs: torch.Tensor, eps: float, x2: torch.Tensor | None = None, y2: torch.Tensor | None = None, x0: torch.Tensor | None = None, max_iter: int = 50, rtol: float = 1e-06, atol: float = 1e-06, autotune: bool = False) -> tuple[torch.Tensor, PythonBatchedCgInfo]","Solve H @ x = b using CG with external apply_plan_vec kernels.","This is a simpler, more reliable implementation that uses the proven\napply_plan_vec_sqeuclid kernels instead of inline Triton softmax.\n\nThe linear operator is:\n    H @ v = denom * v - P^T @ (P @ v \u002F diag_x)\n\nwhere P is the (n, m) transport plan matrix.",[961,963,965,968,971,973,975,977,979,981,983,986,988,990,992],{"name":450,"type":377,"doc":962},"Source points (n, d).",{"name":453,"type":377,"doc":964},"Target points (m, d).",{"name":966,"type":377,"doc":967},"f","Source potential (n,).",{"name":969,"type":377,"doc":970},"g","Target potential (m,).",{"name":614,"type":377,"doc":972},"Diagonal D_x (n,) — row sums of P.",{"name":617,"type":377,"doc":974},"Denominator D_y (m,) — column sums of P plus regularization.",{"name":620,"type":377,"doc":976},"Right-hand side b (m,).",{"name":590,"type":457,"doc":978},"Regularization parameter.",{"name":593,"type":477,"default":468,"doc":980},"Precomputed ``||x||²`` (n,); computed on the fly if None.",{"name":596,"type":477,"default":468,"doc":982},"Precomputed ``||y||²`` (m,); computed on the fly if None.",{"name":984,"type":477,"default":468,"doc":985},"x0","Initial CG guess (m,); zero-initialized if None.",{"name":623,"type":497,"default":987,"doc":625},"50",{"name":627,"type":457,"default":628,"doc":989},"Relative convergence tolerance.",{"name":631,"type":457,"default":628,"doc":991},"Absolute convergence tolerance.",{"name":993,"type":462,"default":463,"doc":994},"autotune","Whether to enable Triton autotuning for apply_plan_vec.",{"type":996,"doc":635},"sol",{"file":922,"line":998},38,{"id":1000,"kind":199,"path":1001,"signature":1003,"summary":1004,"description":1005,"source":1006,"parent":915},"torchmatch.transport.samples.kernels.cg_python_batched.compiled_cg_solve",[177,289,436,508,917,1002],"compiled_cg_solve","def compiled_cg_solve(x: torch.Tensor, y: torch.Tensor, f: torch.Tensor, g: torch.Tensor, diag_x: torch.Tensor, denom: torch.Tensor, rhs: torch.Tensor, eps: float, x2: torch.Tensor | None = None, y2: torch.Tensor | None = None, x0: torch.Tensor | None = None, max_iter: int = 50, rtol: float = 1e-06, atol: float = 1e-06, use_compile: bool = True) -> tuple[torch.Tensor, PythonBatchedCgInfo]","Solve H @ x = b using CG with torch.compile optimization.","This version uses torch.compile for the core CG loop, providing\nCUDA graph capture and kernel fusion for additional speedup.\n\nNote: This runs a fixed number of iterations (no early exit) to be\ntorch.compile friendly. The convergence info reflects the state\nafter max_iter iterations.\n\nArgs:\n    x: Source points (n, d)\n    y: Target points (m, d)\n    f: Source potential (n,)\n    g: Target potential (m,)\n    diag_x: Diagonal D_x (n,) - row sums of P\n    denom: Denominator D_y (m,) - column sums of P plus regularization\n    rhs: Right-hand side b (m,)\n    eps: Regularization parameter\n    x2: Precomputed ||x||^2 (n,) [optional]\n    y2: Precomputed ||y||^2 (m,) [optional]\n    x0: Initial guess (m,) [optional, ignored in compiled version]\n    max_iter: Number of CG iterations to run (fixed, no early exit)\n    rtol, atol: Convergence tolerances (for info only, no early exit)\n    use_compile: Whether to use torch.compile (default True)\n\nReturns:\n    sol: Solution x (m,)\n    info: Convergence information",{"file":922,"line":1007},340,{"id":1009,"kind":181,"path":1010,"signature":1012,"source":1013,"parent":506},"torchmatch.transport.samples.kernels.grad_sqeuclid",[177,289,436,508,1011],"grad_sqeuclid","module torchmatch.transport.samples.kernels.grad_sqeuclid",{"file":1014,"line":187},"torchmatch\u002Ftransport\u002Fsamples\u002Fkernels\u002Fgrad_sqeuclid.py",{"id":1016,"kind":333,"path":1017,"signature":1019,"source":1020,"parent":1009},"torchmatch.transport.samples.kernels.grad_sqeuclid.SHARED_MEM_BUDGET_KB",[177,289,436,508,1011,1018],"SHARED_MEM_BUDGET_KB","SHARED_MEM_BUDGET_KB = 96",{"file":1014,"line":1021},883,{"id":1023,"kind":199,"path":1024,"signature":1026,"source":1027,"parent":1009},"torchmatch.transport.samples.kernels.grad_sqeuclid.sinkhorn_online_grad_sqeuclid",[177,289,436,508,1011,1025],"sinkhorn_online_grad_sqeuclid","def sinkhorn_online_grad_sqeuclid(x: torch.Tensor, y: torch.Tensor, a: torch.Tensor, b: torch.Tensor, f: torch.Tensor, g: torch.Tensor, eps: float, allow_tf32: bool = True, use_exp2: bool = True, block_m: int | None = None, block_n: int | None = None, block_k: int | None = None, num_warps: int | None = None, num_stages: int = 2, autotune: bool = True, grad_scale: torch.Tensor | None = None, compute_grad_x: bool = True, compute_grad_y: bool = True, cost_scale: float = 1.0, label_x: torch.Tensor | None = None, label_y: torch.Tensor | None = None, label_cost_matrix: torch.Tensor | None = None, lambda_x: float = 1.0, lambda_y: float = 0.0) -> tuple[torch.Tensor, torch.Tensor]",{"file":1014,"line":1028},941,{"id":1030,"kind":181,"path":1031,"signature":1033,"summary":1034,"description":1035,"source":1036,"parent":506},"torchmatch.transport.samples.kernels.apply_fused_sqeuclid",[177,289,436,508,1032],"apply_fused_sqeuclid","module torchmatch.transport.samples.kernels.apply_fused_sqeuclid","Fused Schur complement matvec kernel for HVP CG acceleration.","This module implements a two-phase persistent kernel that fuses the axis1 and axis0\ntransport plan applications used in the HVP's CG linear operator, reducing kernel\nlaunches from 2 to 1 per CG iteration.\n\nThe key innovation is using a spin-wait barrier with atomic counter for grid-wide\nsynchronization between phases, allowing both operations to run in a single kernel.\n\nIMPORTANT: The spin barrier requires all blocks to be co-resident on the GPU.\nIf grid_size exceeds the maximum resident blocks (num_SMs * ~4), the kernel\nwill deadlock. This module automatically falls back to the two-kernel approach\nwhen the grid size would be too large.",{"file":1037,"line":187},"torchmatch\u002Ftransport\u002Fsamples\u002Fkernels\u002Fapply_fused_sqeuclid.py",{"id":1039,"kind":199,"path":1040,"signature":1042,"summary":1043,"description":1044,"source":1045,"parent":1030},"torchmatch.transport.samples.kernels.apply_fused_sqeuclid.fused_schur_matvec_sqeuclid",[177,289,436,508,1032,1041],"fused_schur_matvec_sqeuclid","def fused_schur_matvec_sqeuclid(x: torch.Tensor, y: torch.Tensor, f: torch.Tensor, g: torch.Tensor, z: torch.Tensor, diag_x: torch.Tensor, denom: torch.Tensor, eps: float, x2: torch.Tensor | None = None, y2: torch.Tensor | None = None, piz_buffer: torch.Tensor | None = None, block_m: int = 64, block_n: int = 64, block_k: int = 64, num_warps: int = 4, num_stages: int = 2, use_exp2: bool = True, allow_tf32: bool = False) -> tuple[torch.Tensor, torch.Tensor]","Fused Schur complement matvec: out = denom * z - P^T @ (P @ z \u002F diag_x).","This is a single-kernel implementation that fuses two transport plan applications\n(axis1 and axis0) with a grid-wide barrier for synchronization.\n\nArgs:\n    x: Source points (n, d)\n    y: Target points (m, d)\n    f: Source potential (n,) - raw-form: P = exp((f+g-C)\u002Feps)\n    g: Target potential (m,) - raw-form: P = exp((f+g-C)\u002Feps)\n    z: Input CG vector (m,)\n    diag_x: Source diagonal D_x = diag_factor_x * a_hat (n,)\n    denom: Target denominator D_y = diag_factor_y * b_hat + eps*tau2 (m,)\n    eps: Entropy regularization\n    x2: Precomputed ||x||^2 (n,) [optional]\n    y2: Precomputed ||y||^2 (m,) [optional]\n    piz_buffer: Reusable buffer for P @ z intermediate (n,) [optional]\n\nReturns:\n    out: Result denom * z - P^T @ (P @ z \u002F diag_x), shape (m,)\n    piz_buffer: The intermediate buffer (for reuse)\n\nNote:\n    This kernel uses a spin-wait barrier for grid-wide synchronization.\n    The grid size must not exceed the number of co-resident blocks on the GPU.\n    If the grid would be too large, this function automatically falls back to\n    a two-kernel approach using apply_plan_vec_shifted (axis=1, then axis=0).\n\n    The fallback path converts raw potentials to shifted form and uses shifted-form apply kernels.",{"file":1037,"line":1046},338,{"id":1048,"kind":181,"path":1049,"signature":1051,"summary":1052,"description":1053,"source":1054,"parent":506},"torchmatch.transport.samples.kernels.apply_raw",[177,289,436,508,1050],"apply_raw","module torchmatch.transport.samples.kernels.apply_raw","raw-form apply kernels (mat5 for HVP, deprecated vec\u002Fmat wrappers).","This module contains:\n- mat5_sqeuclid: Active kernel for HVP computation (P_ij * (A_i . y_j) * y_j)\n- apply_plan_vec_sqeuclid: Deprecated, delegates to shifted-form\n- apply_plan_mat_sqeuclid: Deprecated, delegates to shifted-form\n\nThe mat5 kernel uses raw-form potentials (f, g include absorbed log marginals).\nThe deprecated wrappers convert raw potentials to shifted form and call shifted-form.",{"file":1055,"line":187},"torchmatch\u002Ftransport\u002Fsamples\u002Fkernels\u002Fapply_raw.py",{"id":1057,"kind":199,"path":1058,"signature":1060,"summary":1061,"description":1062,"source":1063,"parent":1048},"torchmatch.transport.samples.kernels.apply_raw.apply_plan_vec_sqeuclid",[177,289,436,508,1050,1059],"apply_plan_vec_sqeuclid","def apply_plan_vec_sqeuclid(x: torch.Tensor, y: torch.Tensor, f: torch.Tensor, g: torch.Tensor, vec: torch.Tensor, eps: float, axis: int, cost_scale: float = 1.0, x2: torch.Tensor | None = None, y2: torch.Tensor | None = None, log_a: torch.Tensor | None = None, log_b: torch.Tensor | None = None, block_m: int | None = None, block_n: int | None = None, block_k: int | None = None, num_warps: int = 4, num_stages: int = 2, use_exp2: bool = True, allow_tf32: bool = False, autotune: bool = True) -> torch.Tensor","Apply P or Pt to a vector without materializing P (streaming, stable).",".. deprecated::\n    Use :func:`apply_plan_vec_shifted` instead for better performance\n    with shifted-form shifted potentials.\n\nComputes:\n  axis=1: out[i] = sum_j exp((f_i + g_j - cost_scale*C_ij)\u002Feps) * vec[j]\n  axis=0: out[j] = sum_i exp((f_i + g_j - cost_scale*C_ij)\u002Feps) * vec[i]\n\nArgs:\n    f: raw-form source potential [n] (includes absorbed log marginal)\n    g: raw-form target potential [m] (includes absorbed log marginal)\n    log_a: Optional log source weights [n]. If None, assumes uniform (log(1\u002Fn)).\n    log_b: Optional log target weights [m]. If None, assumes uniform (log(1\u002Fm)).\n    cost_scale: Scaling for cost function. 1.0 for full ||x-y||^2, 0.5 for half.\n    autotune: If True (default), use autotuned kernel configs for best performance.\n              If False, use manual block sizes (useful for reproducible benchmarks).",{"file":1055,"line":1064},221,{"id":1066,"kind":199,"path":1067,"signature":1069,"summary":1070,"description":1071,"source":1072,"parent":1048},"torchmatch.transport.samples.kernels.apply_raw.apply_plan_mat_sqeuclid",[177,289,436,508,1050,1068],"apply_plan_mat_sqeuclid","def apply_plan_mat_sqeuclid(x: torch.Tensor, y: torch.Tensor, f: torch.Tensor, g: torch.Tensor, mat: torch.Tensor, eps: float, axis: int, cost_scale: float = 1.0, scale: torch.Tensor | None = None, x2: torch.Tensor | None = None, y2: torch.Tensor | None = None, log_a: torch.Tensor | None = None, log_b: torch.Tensor | None = None, block_m: int | None = None, block_n: int | None = None, block_k: int | None = None, block_d: int | None = None, num_warps: int = 4, num_stages: int = 2, use_exp2: bool = True, allow_tf32: bool = False, autotune: bool = True) -> torch.Tensor","Apply P or Pt to a matrix without materializing P (streaming, stable).",".. deprecated::\n    Use :func:`apply_plan_mat_shifted` instead for better performance\n    with shifted-form shifted potentials.\n\nComputes:\n  axis=1: out[i, :] = sum_j exp((f_i + g_j - cost_scale*C_ij)\u002Feps) * mat[j, :]\n  axis=0: out[j, :] = sum_i exp((f_i + g_j - cost_scale*C_ij)\u002Feps) * mat[i, :]\n\nIf ``scale`` is provided (axis=1 only), the kernel uses `mat[j,:] * scale[j]`\non the fly (avoids allocating `mat * scale[:,None]`).\n\nArgs:\n    f: raw-form source potential [n] (includes absorbed log marginal)\n    g: raw-form target potential [m] (includes absorbed log marginal)\n    log_a: Optional log source weights [n]. If None, assumes uniform (log(1\u002Fn)).\n    log_b: Optional log target weights [m]. If None, assumes uniform (log(1\u002Fm)).\n    cost_scale: Scaling for cost function. 1.0 for full ||x-y||^2, 0.5 for half.\n    autotune: If True (default), use autotuned kernel configs for best performance.\n              If False, use manual block sizes (useful for reproducible benchmarks).",{"file":1055,"line":1073},342,{"id":1075,"kind":199,"path":1076,"signature":1078,"summary":1079,"description":1080,"source":1081,"parent":1048},"torchmatch.transport.samples.kernels.apply_raw.mat5_sqeuclid",[177,289,436,508,1050,1077],"mat5_sqeuclid","def mat5_sqeuclid(x: torch.Tensor, y: torch.Tensor, f: torch.Tensor, g: torch.Tensor, A: torch.Tensor, eps: float, cost_scale: float = 1.0, x2: torch.Tensor | None = None, y2: torch.Tensor | None = None, block_m: int | None = None, block_n: int | None = None, block_k: int | None = None, num_warps: int = 4, num_stages: int = 2, use_exp2: bool = True, allow_tf32: bool = False, autotune: bool = True) -> torch.Tensor","Compute Mat5 term for HVP: Mat5 = (-4*cost_scale\u002Feps) * sum_j P_ij (A_i.y_j) y_j.","Args:\n    autotune: If True (default), use autotuned kernel configs for best performance.\n              If False, use manual block sizes (useful for reproducible benchmarks).",{"file":1055,"line":1082},492,{"id":1084,"kind":181,"path":1085,"signature":1087,"summary":1088,"description":1089,"source":1090,"parent":506},"torchmatch.transport.samples.kernels.c_transform_sqeuclid",[177,289,436,508,1086],"c_transform_sqeuclid","module torchmatch.transport.samples.kernels.c_transform_sqeuclid","C-Transform (hard argmin) kernel for squared Euclidean cost.","Computes the non-entropic Kantorovich c-transform via streaming min + argmin:\n\n    c_i = min_j [cost_scale * ||x_i - y_j||² - ψ_j]\n    j*_i = argmin_j [cost_scale * ||x_i - y_j||² - ψ_j]\n\nFactorization (same trick as the LSE kernel):\n\n    cost_scale * ||x-y||² - ψ = cost_scale*||x||² + (cost_scale*||y||² - ψ) - 2*cost_scale*(x·y)\n                               = alpha_i + bias_j - coord_scale * dot(x_i, y_j)\n\nwhere:\n    alpha_i = cost_scale * ||x_i||²  (constant w.r.t. j, factors out of min)\n    bias_j  = cost_scale * ||y_j||² - ψ_j\n    coord_scale = 2 * cost_scale\n\nThe kernel computes min_j[-coord_scale * dot(x_i, y_j) + bias_j] via tiled streaming.\nThe Python wrapper adds alpha_i back to get the final c-transform values.\n\nTie-breaking: across tiles, smallest-j wins (strict \u003C comparison).\nWithin a tile, tl.argmin selects the first minimum per Triton lane ordering.\n\nKernel outputs int32 indices (saves SRAM\u002Fregisters). Python wrapper casts to int64.",{"file":1091,"line":187},"torchmatch\u002Ftransport\u002Fsamples\u002Fkernels\u002Fc_transform_sqeuclid.py",{"id":1093,"kind":199,"path":1094,"signature":1096,"summary":1097,"description":1098,"source":1099,"parent":1084},"torchmatch.transport.samples.kernels.c_transform_sqeuclid.c_transform_kernel",[177,289,436,508,1086,1095],"c_transform_kernel","def c_transform_kernel(x: torch.Tensor, y: torch.Tensor, bias: torch.Tensor, cost_scale: float = 1.0, allow_tf32: bool = True, block_m: int | None = None, block_n: int | None = None, block_k: int | None = None, num_warps: int | None = None, num_stages: int = 3, autotune: bool = True) -> tuple[torch.Tensor, torch.Tensor]","Compute streaming min + argmin over the factored cost.","Computes per source point i:\n    min_val_i = min_j [-coord_scale * dot(x_i, y_j) + bias_j]\n    min_idx_i = argmin_j [same]\n\nThis is the inner minimum only. The caller adds alpha_i = cost_scale * ||x_i||²\nto get the full c-transform values.\n\nArgs:\n    x: Source coordinates [n, d], CUDA\n    y: Target coordinates [m, d], CUDA\n    bias: Pre-scaled bias [m] = cost_scale * ||y||² - ψ\n    cost_scale: Cost scaling (1.0 for ||x-y||², 0.5 for ||x-y||²\u002F2)\n    allow_tf32: Enable TF32 for matmul\n    block_m, block_n, block_k: Manual block sizes (disables autotune)\n    num_warps: Number of warps (disables autotune)\n    num_stages: Pipeline stages\n    autotune: Enable autotuning\n\nReturns:\n    min_vals: Inner minimum values [n], float32\n    argmin_idx: Argmin indices [n], int64 (cast from kernel int32)",{"file":1091,"line":1100},239,{"id":1102,"kind":181,"path":1103,"signature":1105,"summary":1106,"description":1107,"source":1108,"parent":506},"torchmatch.transport.samples.kernels.apply_shifted",[177,289,436,508,1104],"apply_shifted","module torchmatch.transport.samples.kernels.apply_shifted","shifted-form apply kernels (shifted potentials, s_I cancellation).","This module contains streaming P @ V and P @ vec kernels that work with\nSHIFTED potentials directly (f_hat = f - alpha, g_hat = g - beta), avoiding\nthe cost of converting between potential conventions.\n\nKernels:\n- apply_plan_mat_shifted: P @ mat or P^T @ mat (2D grid, tiles over D)\n- apply_plan_vec_shifted: P @ vec or P^T @ vec (1D grid, s_I cancellation)\n\nKey insight: Compute score WITHOUT f_hat\u002Fg_hat in the tiled loop, then apply\nrow\u002Fcolumn marginal correction at the end. This is the same numerical trick as\nthe online-softmax correction factor used in streaming attention kernels.",{"file":1109,"line":187},"torchmatch\u002Ftransport\u002Fsamples\u002Fkernels\u002Fapply_shifted.py",{"id":1111,"kind":199,"path":1112,"signature":1114,"summary":1115,"description":1116,"source":1117,"parent":1102},"torchmatch.transport.samples.kernels.apply_shifted.apply_plan_mat_shifted",[177,289,436,508,1104,1113],"apply_plan_mat_shifted","def apply_plan_mat_shifted(x: torch.Tensor, y: torch.Tensor, f_hat: torch.Tensor, g_hat: torch.Tensor, log_a: torch.Tensor, log_b: torch.Tensor, mat: torch.Tensor, eps: float, axis: int, cost_scale: float = 1.0, scale: torch.Tensor | None = None, block_m: int | None = None, block_n: int | None = None, block_k: int | None = None, block_d: int | None = None, num_warps: int = 4, num_stages: int = 2, use_exp2: bool = True, allow_tf32: bool = False, autotune: bool = True) -> torch.Tensor","Apply P or P^T to a matrix using shifted potentials.","This kernel works with SHIFTED potentials directly:\n    f_hat = f - alpha  where alpha = cost_scale * ||x||^2\n    g_hat = g - beta   where beta  = cost_scale * ||y||^2\n\nComputes:\n  axis=1: out[i, :] = sum_j P_ij * mat[j, :]  (P @ V)\n  axis=0: out[j, :] = sum_i P_ij * mat[i, :]  (P^T @ V)\n\nwhere P_ij = a_i * b_j * exp((f_hat + g_hat + 2*cost_scale*x.y) \u002F eps)\n\nArgs:\n    x: Source points [n, d]\n    y: Target points [m, d]\n    f_hat: Shifted f potential [n] (f_hat = f - cost_scale * ||x||^2)\n    g_hat: Shifted g potential [m] (g_hat = g - cost_scale * ||y||^2)\n    log_a: Log source weights [n]\n    log_b: Log target weights [m]\n    mat: Matrix V [m, d] for axis=1, [n, d] for axis=0\n    eps: Regularization parameter\n    axis: 1 for P @ V, 0 for P^T @ V\n    cost_scale: Scaling for cost (1.0 for ||x-y||^2, 0.5 for ||x-y||^2\u002F2)\n    scale: Optional per-row scale [m] for axis=1 (applies mat * scale[:, None])\n    autotune: If True (default), use autotuned kernel configs\n\nReturns:\n    out: Result matrix [n, d] for axis=1, [m, d] for axis=0\n\nNote:\n    The shifted potentials can be obtained from standard potentials via:\n        f_hat = f - cost_scale * (x ** 2).sum(dim=1)\n        g_hat = g - cost_scale * (y ** 2).sum(dim=1)\n\n    Or directly from the shifted-potential Sinkhorn solvers.",{"file":1109,"line":685},{"id":1119,"kind":199,"path":1120,"signature":1122,"summary":1123,"description":1124,"source":1125,"parent":1102},"torchmatch.transport.samples.kernels.apply_shifted.apply_plan_vec_shifted",[177,289,436,508,1104,1121],"apply_plan_vec_shifted","def apply_plan_vec_shifted(x: torch.Tensor, y: torch.Tensor, f_hat: torch.Tensor, g_hat: torch.Tensor, log_a: torch.Tensor, log_b: torch.Tensor, vec: torch.Tensor, eps: float, axis: int, cost_scale: float = 1.0, allow_tf32: bool = False, use_exp2: bool = True, block_m: int = 64, block_n: int = 64, block_k: int = 64, num_warps: int = 4, num_stages: int = 2) -> torch.Tensor","Apply transport plan to vector using shifted-form with s_I cancellation.","Computes P @ vec (axis=1) or P^T @ vec (axis=0) where:\nP[i,j] = a[i] * b[j] * exp((f_hat[i] + g_hat[j] + 2*cost_scale*x[i].y[j]) \u002F eps)\n\nKey optimization: The normalizing sum s_I cancels algebraically, so we use:\n- axis=1: out_I = a_I * exp(f_hat_I\u002Feps + m_I) * O_I\n- axis=0: out_J = b_J * exp(g_hat_J\u002Feps + m_J) * O_J\n\nArgs:\n    x: Source points [n, d], fp16\u002Ffp32\n    y: Target points [m, d], fp16\u002Ffp32\n    f_hat: Shifted source potential [n], fp32\n    g_hat: Shifted target potential [m], fp32\n    log_a: Log source weights [n], fp32\n    log_b: Log target weights [m], fp32\n    vec: Vector to apply [m] for axis=1, [n] for axis=0\n    eps: Regularization parameter\n    axis: 1 for P @ vec, 0 for P^T @ vec\n    cost_scale: Scaling for cost (0.5 for half_cost)\n    allow_tf32: Use TF32 for dot products\n    use_exp2: Use exp2 instead of exp\n    block_m, block_n, block_k: Block sizes\n    num_warps, num_stages: Triton tuning params\n\nReturns:\n    out: Result vector [n] for axis=1, [m] for axis=0",{"file":1109,"line":1126},961,{"id":1128,"kind":181,"path":1129,"signature":1131,"summary":1132,"description":1133,"source":1134,"parent":177},"torchmatch.assignment",[177,1130],"assignment","module torchmatch.assignment","Integer linear assignment problem (LAP) solvers.","Public surface:\n\n- :func:`solve` -- the dispatcher; resolves AUTO at call time so the\n  picked op traces cleanly under ``torch.compile``.\n- :class:`Backend` -- backend choices accepted by :func:`solve`.\n- :mod:`ops` -- direct access to ``torch.ops.assignment.*`` op handles\n  for callers that want to pin a backend.",{"file":1135,"line":187},"torchmatch\u002Fassignment\u002F__init__.py",{"id":1137,"kind":199,"path":1138,"signature":299,"summary":1139,"description":1140,"source":1141,"parent":1128},"torchmatch.assignment.load_cpu",[177,1130,298],"Register ``torch.ops.assignment.jonker_*`` (CPU backend).","Prefers a prebuilt extension shipped in the wheel and falls back\nto JIT-compiling the C++ sources via\n:func:`torch.utils.cpp_extension.load`. Set\n``TORCHMATCH_FORCE_JIT=1`` to skip the prebuilt path.",{"file":1142,"line":1143},"torchmatch\u002Fassignment\u002F_cpu.py",64,{"id":1145,"kind":199,"path":1146,"signature":309,"summary":1147,"description":1148,"source":1149,"parent":1128},"torchmatch.assignment.load_cuda",[177,1130,308],"Register the CUDA-only ops in ``torch.ops.assignment``.","Adds Munkres classical (``munkres``), the experimental Munkres\nhybrid (``hybrid``), Lawler tree-augmentation (``lawler``), and the\nCUDA backend of ``jonker_dense_batch``. Prefers a prebuilt extension\nshipped in the wheel and falls back to JIT-compiling the C++\u002FCUDA\nsources via :func:`torch.utils.cpp_extension.load`. Set\n``TORCHMATCH_FORCE_JIT=1`` to skip the prebuilt path.",{"file":1150,"line":1151},"torchmatch\u002Fassignment\u002F_cuda.py",61,{"id":1153,"kind":199,"path":1154,"signature":1156,"summary":1157,"description":1158,"params":1159,"returns":1170,"source":1172,"parent":1128},"torchmatch.assignment.auction_assignment",[177,1130,1155],"auction_assignment","def auction_assignment(cost_matrix: torch.Tensor, bid_size: float, max_iters: int = 100000) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]","Solve a linear assignment problem using Bertsekas' auction algorithm.","Converts the cost matrix to a profit matrix, then runs synchronous\nbidding until every row (or every column, for rectangular problems) is\nassigned. Non-finite entries mark forbidden pairs and are penalised\ninternally so the solver never selects them.",[1160,1163,1166],{"name":1161,"type":377,"doc":1162},"cost_matrix","``(N, M)`` cost matrix. ``+inf`` entries mark forbidden pairs.\nNaN and ``-inf`` are rejected.",{"name":1164,"type":457,"doc":1165},"bid_size","Auction bid step size. The internal ``epsilon`` is derived as\n``min(bid_size \u002F min(N, M), 1e-3)``.",{"name":1167,"type":497,"default":1168,"doc":1169},"max_iters","100000","Maximum number of bidding iterations. Raises ``RuntimeError`` if\nthe algorithm has not converged within this budget.",{"type":377,"doc":1171},"``(K, 2)`` long tensor of matched ``(row, col)`` indices.",{"file":1173,"line":1174},"torchmatch\u002Fassignment\u002F_auction.py",10,{"id":1176,"kind":199,"path":1177,"signature":1179,"summary":1180,"params":1181,"returns":1193,"source":1196,"parent":1128},"torchmatch.assignment.assignment_cost",[177,1130,1178],"assignment_cost","def assignment_cost(cost: torch.Tensor, matches: torch.Tensor, reduction: str = 'sum') -> torch.Tensor","Compute the total cost of a LAP assignment.",[1182,1185,1188],{"name":1183,"type":377,"doc":1184},"cost","Cost matrix (N, M) or (B, N, M). float32 or float64.",{"name":1186,"type":377,"doc":1187},"matches","Row→col assignment (N,) or (B, N). int64. Unmatched rows have ``-1``.",{"name":1189,"type":1190,"default":1191,"doc":1192},"reduction","str","'sum'","How to aggregate per-row costs: ``\"sum\"`` (default) sums all matched\nrows; ``\"mean\"`` divides by the number of matched rows; ``\"none\"``\nreturns per-row costs with unmatched rows set to 0.",{"type":1194,"doc":1195},"total","Scalar for 2-D input or (B,) for 3-D input, unless\n``reduction=\"none\"``, in which case the shape is (N,) or (B, N).",{"file":1197,"line":1198},"torchmatch\u002Fassignment\u002F_cost.py",8,{"id":1200,"kind":324,"path":1201,"signature":327,"source":1202,"parent":1128},"torchmatch.assignment.Backend",[177,1130,326],{"file":1203,"line":1204},"torchmatch\u002Fassignment\u002F_solve.py",25,{"id":1206,"kind":333,"path":1207,"signature":336,"source":1208,"parent":1200},"torchmatch.assignment.Backend.AUTO",[177,1130,326,335],{"file":1203,"line":1209},26,{"id":1211,"kind":333,"path":1212,"signature":1214,"source":1215,"parent":1200},"torchmatch.assignment.Backend.JONKER",[177,1130,326,1213],"JONKER","JONKER = 'jonker'",{"file":1203,"line":1216},27,{"id":1218,"kind":333,"path":1219,"signature":1221,"source":1222,"parent":1200},"torchmatch.assignment.Backend.MUNKRES",[177,1130,326,1220],"MUNKRES","MUNKRES = 'munkres'",{"file":1203,"line":934},{"id":1224,"kind":333,"path":1225,"signature":1227,"source":1228,"parent":1200},"torchmatch.assignment.Backend.LAWLER",[177,1130,326,1226],"LAWLER","LAWLER = 'lawler'",{"file":1203,"line":432},{"id":1230,"kind":333,"path":1231,"signature":1233,"source":1234,"parent":1200},"torchmatch.assignment.Backend.GREEDY",[177,1130,326,1232],"GREEDY","GREEDY = 'greedy'",{"file":1203,"line":943},{"id":1236,"kind":199,"path":1237,"signature":1239,"summary":1240,"description":1241,"params":1242,"returns":1250,"source":1253,"parent":1128},"torchmatch.assignment.resolve_backend",[177,1130,1238],"resolve_backend","def resolve_backend(cost: torch.Tensor, backend: Backend | str = Backend.AUTO) -> str","Return the op name that ``solve()`` would dispatch to.","Useful for debugging AUTO routing and for asserting which backend fires\nin a test or benchmark.",[1243,1245],{"name":1183,"type":377,"doc":1244},"Cost matrix (N, M) or (B, N, M) on the target device.",{"name":1246,"type":1247,"default":1248,"doc":1249},"backend","Backend | str","Backend.AUTO","Backend hint. Non-``AUTO`` values are echoed back as their string value\n(e.g. ``\"munkres\"``).",{"type":1251,"doc":1252},"op_name","Exact ``torch.ops.assignment.\u003Cname>`` key — e.g. ``\"jonker_compact\"``,\n``\"munkres\"``, ``\"lawler\"``, ``\"greedy\"``.",{"file":1203,"line":1254},291,{"id":1256,"kind":199,"path":1257,"signature":1258,"source":1259,"parent":1128},"torchmatch.assignment.solve",[177,1130,393],"def solve(cost: torch.Tensor, backend: Backend | str = Backend.AUTO, unpack: bool = False) -> torch.Tensor | tuple[torch.Tensor, ...]",{"file":1203,"line":1260},265,{"id":1262,"kind":181,"path":1263,"signature":1264,"summary":1265,"description":1266,"source":1267,"parent":1128},"torchmatch.assignment.ops",[177,1130,400],"module torchmatch.assignment.ops","Direct access to ``torch.ops.assignment.*`` ops.","Each attribute is the corresponding op handle, so callers that want to\npin a specific backend (e.g. for benchmarking or to sidestep the\ndispatcher's AUTO branch) write::\n\n    from torchmatch.assignment.ops import jonker_dense\n\n    result = jonker_dense(cost)",{"file":1268,"line":187},"torchmatch\u002Fassignment\u002Fops.py",{"id":1270,"kind":199,"path":1271,"signature":1273,"source":1274,"parent":1262},"torchmatch.assignment.ops.jonker_scalar",[177,1130,400,1272],"jonker_scalar","def jonker_scalar(cost: torch.Tensor) -> torch.Tensor",{"file":1268,"line":411},{"id":1276,"kind":199,"path":1277,"signature":1279,"source":1280,"parent":1262},"torchmatch.assignment.ops.jonker_dense",[177,1130,400,1278],"jonker_dense","def jonker_dense(cost: torch.Tensor) -> torch.Tensor",{"file":1268,"line":1281},4,{"id":1283,"kind":199,"path":1284,"signature":1286,"source":1287,"parent":1262},"torchmatch.assignment.ops.jonker_compact",[177,1130,400,1285],"jonker_compact","def jonker_compact(cost: torch.Tensor) -> torch.Tensor",{"file":1268,"line":1288},5,{"id":1290,"kind":199,"path":1291,"signature":1293,"source":1294,"parent":1262},"torchmatch.assignment.ops.jonker_dense_batch",[177,1130,400,1292],"jonker_dense_batch","def jonker_dense_batch(cost: torch.Tensor) -> torch.Tensor",{"file":1268,"line":1295},6,{"id":1297,"kind":199,"path":1298,"signature":1300,"source":1301,"parent":1262},"torchmatch.assignment.ops.jonker_compact_batch",[177,1130,400,1299],"jonker_compact_batch","def jonker_compact_batch(cost: torch.Tensor) -> torch.Tensor",{"file":1268,"line":1302},7,{"id":1304,"kind":199,"path":1305,"signature":1307,"source":1308,"parent":1262},"torchmatch.assignment.ops.jonker_dense_batch_unpacked",[177,1130,400,1306],"jonker_dense_batch_unpacked","def jonker_dense_batch_unpacked(cost: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]",{"file":1268,"line":1198},{"id":1310,"kind":199,"path":1311,"signature":1313,"source":1314,"parent":1262},"torchmatch.assignment.ops.jonker_compact_batch_unpacked",[177,1130,400,1312],"jonker_compact_batch_unpacked","def jonker_compact_batch_unpacked(cost: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]",{"file":1268,"line":1315},11,{"id":1317,"kind":199,"path":1318,"signature":1320,"source":1321,"parent":1262},"torchmatch.assignment.ops.greedy",[177,1130,400,1319],"greedy","def greedy(cost: torch.Tensor) -> torch.Tensor",{"file":1268,"line":1322},20,{"id":1324,"kind":199,"path":1325,"signature":1327,"source":1328,"parent":1262},"torchmatch.assignment.ops.munkres",[177,1130,400,1326],"munkres","def munkres(cost: torch.Tensor) -> torch.Tensor",{"file":1268,"line":1329},17,{"id":1331,"kind":199,"path":1332,"signature":1334,"source":1335,"parent":1262},"torchmatch.assignment.ops.hybrid",[177,1130,400,1333],"hybrid","def hybrid(cost: torch.Tensor) -> torch.Tensor",{"file":1268,"line":425},{"id":1337,"kind":199,"path":1338,"signature":1340,"source":1341,"parent":1262},"torchmatch.assignment.ops.lawler",[177,1130,400,1339],"lawler","def lawler(cost: torch.Tensor) -> torch.Tensor",{"file":1268,"line":1342},19,null,1785218163492]