Skip to content

[PyTorch] Avoid concatenation copies for compiled split Linear parameters - #3623

Draft
pggPL wants to merge 4 commits into
NVIDIA:mainfrom
pggPL:split_parameters_linear
Draft

pggPL wants to merge 4 commits into
NVIDIA:mainfrom
pggPL:split_parameters_linear

Conversation

@pggPL

@pggPL pggPL commented Oct 5, 2026 •

Copy link
Copy Markdown
Collaborator

Description

Compiled te.Linear(parameters_split=...) can consume adjacent parameter parts
without a concatenation copy. Pass the original parts through the custom op and
construct a checked view inside the consumer, preserving separate Parameters,
state-dict names and gradients. Disjoint or noncontiguous parts use concatenation.

This PR targets upstream main and adds split-parameter compile support only to
te.Linear. ConcatInput[T] describes an ordinary operand or immutable
ParameterParts. A single concat_input helper keeps singleton tensors unchanged,
uses noop_cat in eager, and defers supported compiled operands. Original parts
cross the op boundary; fake implementations receive metadata and the real
consumer constructs the view. Linear retains its existing final configuration
validation and setup_saved_tensors(ctx) hook.

Type of change

  • Non-breaking performance improvement and bug fix

Changes

  • Make the early fallback check inspect original parameters without
    concatenating them; retain final validation of the prepared op arguments.
  • Save original parts through autograd and split full gradients outside the
    backward op; recheck storage adjacency on each execution.
  • Compute bias gradients when frozen weights skip the weight-gradient GEMM.
  • Preserve dtype promotion under autocast and copy lazy negative/conjugate views.
  • Keep returned or unfused split biases as traceable concatenations under compile.

Validation and remaining work

  • RTX 5880 Ada, PyTorch 2.15 development build/CUDA 13.4, compatible prebuilt
    native TE extension: complete test_torch_compile.py passed (187 passed,
    46 skipped, 1 existing non-strict xfail passed). Includes 45 focused split/API
    cases covering storage and parameter replacement, CUDA Graphs, saved hooks,
    version checks, gradient targets, recipes and returned bias.
  • Pre-commit, Black, full Python L0 lint (10.00/10), license and whitespace checks
    passed; tested source hashes match the commit.
  • Performance work remains: disjoint storage can be concatenated repeatedly in
    backward, and earlier split-projection measurements found host overhead for
    small shapes. Ordinary-Linear CPU latency has not been benchmarked. No general
    speedup or Blackwell/distributed validation is claimed.

Checklist

  • Added tests and updated the Linear parameter documentation
  • Run full Python L0 lint
  • Run GPU regression tests on this revision
  • Complete performance validation and human review before marking ready

pggPL added 4 commits October 5, 2026 11:42
…ters

Pass split Linear parameters through the custom-op boundary, reconstruct checked views at execution time and split full gradients outside backward. Preserve concatenation for disjoint storage and returned biases, handle promoted dtypes, and compute bias gradients when weights are frozen.

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Replace ConcatenatedTensor with DeferredCat(parts), expose metadata through to_spec and keep materialization inside the custom op. Preserve deferred backward operands for non-FP8 training, reject quantized parts, and make dimension checks consume explicit metadata.

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Keep singleton operands unwrapped and represent split parameters with immutable ParameterParts through the ConcatInput alias. Centralize eager and compiled input preparation, let the custom-op adapter restore saved parts, and select the execution path before concatenating parameters. Name backward weights by their role and update the existing MLA caller.

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Restore compile_unsupported_reason, setup_saved_tensors(ctx), the original backward operand names, and the existing adapter contract. Remove the unrelated MLA changes. Retain the parameter-parts alias and helper, split-parameter correctness fixes, and focused regression coverage.

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant