test: add comprehensive DistributedTraining integration tests - #754
Conversation
Add 47 integration tests for the DistributedTraining module covering: - InMemoryCommunicationBackend operations (AllReduce, AllGather, Broadcast, Scatter, Barrier, Send/Receive) - DDPModel, FSDPModel, ZeRO-1/2/3 Model tests - PipelineParallelModel, TensorParallelModel, HybridShardedModel tests - ShardingConfiguration validation and edge cases Bug fixes discovered through testing: - Fix DivideByZeroException in PipelineParallelModel.InitializeSharding() - Fix DivideByZeroException in TensorParallelModel.InitializeSharding() - Fix ArgumentException in HybridShardedModel constructor Root cause: C# constructor initialization order - derived class fields were uninitialized when base constructor called virtual InitializeSharding(). Solution: Added OnBeforeInitializeSharding() virtual method to ShardedModelBase allowing derived classes to set up state before sharding initialization. Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>
|
The latest updates on your projects. Learn more about Vercel for GitHub.
|
Summary by CodeRabbit
✏️ Tip: You can customize this high-level summary in your review settings. WalkthroughThis PR introduces a deferred initialization hook ( Changes
Sequence Diagram(s)sequenceDiagram
actor DC as Derived Class<br/>(HybridShardedModel)
participant PendingConfig as ThreadLocal<br/>PendingConfig
participant Base as ShardedModelBase<br/>Constructor
participant Hook as OnBeforeInitialize<br/>Sharding()
participant InitShard as InitializeSharding()
DC->>PendingConfig: StoreConfigAndPassThrough<br/>(pipeline, tensor, data sizes)
PendingConfig-->>DC: Store values in ThreadLocal
DC->>Base: Call base constructor
Base->>Hook: Call OnBeforeInitializeSharding()
Hook->>PendingConfig: Read deferred config
PendingConfig-->>Hook: Return stored values
Hook->>Hook: Initialize fields from<br/>PendingConfig values
Base->>InitShard: Call InitializeSharding()
InitShard->>InitShard: Use initialized fields<br/>for sharding logic
Base-->>DC: Constructor completes
DC->>PendingConfig: Clear ThreadLocal
Estimated code review effort🎯 4 (Complex) | ⏱️ ~60 minutes Possibly related PRs
Suggested labels
Poem
🚥 Pre-merge checks | ✅ 4 | ❌ 1❌ Failed checks (1 warning)
✅ Passed checks (4 passed)
✏️ Tip: You can configure your own custom pre-merge checks in the settings. ✨ Finishing touches
🧪 Generate unit tests (beta)
Comment |
| CachedFullParameters = null; | ||
|
|
||
| // Allow derived classes to set up state before sharding | ||
| OnBeforeInitializeSharding(); |
Check warning
Code scanning / CodeQL
Virtual call in constructor or destructor Warning
Show autofix suggestion
Hide autofix suggestion
Copilot Autofix
AI 9 months ago
In general, to fix a “virtual call in constructor” problem, you move the virtual call out of the constructor and into a separate initialization step that is invoked after the object is fully constructed, or you replace the virtual hook with a non‑virtual callback (e.g., passing a delegate or using a factory). For this class, we should ensure that OnBeforeInitializeSharding() is not invoked from the base constructor while still preserving the intended initialization sequence: (1) base state, (2) subclass custom pre‑sharding logic, (3) sharding.
The least intrusive and safest approach, given we can only edit this file, is:
- Introduce a new non‑virtual instance method (e.g.,
Initialize()) that performs the current constructor’s “late” work: it callsOnBeforeInitializeSharding()andInitializeSharding(). This method can remainprotectedso that external callers don’t depend on it unless they choose to. - Remove the direct calls to
OnBeforeInitializeSharding()andInitializeSharding()from the constructor so that the base constructor only initializes base fields and the communication backend. - Document (via comments) that derived classes are expected to invoke
Initialize()after construction, or that some external factory/creator will do so. Since we cannot edit other files, we must avoid assuming or changing external call sites; instead, we preserve behavior by providing the new method but ensure the constructor no longer makes the problematic virtual call.
Within src/DistributedTraining/ShardedModelBase.cs, this means:
- Keep the constructor’s field initialization (lines 121–135) as‑is.
- Remove lines 137–140 that call
OnBeforeInitializeSharding(); InitializeSharding();. - Add a new
protected void Initialize()method (or similarly named) right after the constructor that performs those two calls in the same order. This method can safely be called by derived classes at the end of their constructors when they are fully initialized.
No new imports or external dependencies are required; the change is purely structural within this class.
| @@ -128,12 +128,29 @@ | ||
| Config.CommunicationBackend.Initialize(); | ||
| } | ||
|
|
||
| // Initialize sharding | ||
| // Initialize sharding-related fields with safe defaults. | ||
| LocalShard = new Vector<T>(Array.Empty<T>()); | ||
| ShardStartIndex = 0; | ||
| ShardSize = 0; | ||
| CachedFullParameters = null; | ||
|
|
||
| // NOTE: | ||
| // Do not call virtual methods such as OnBeforeInitializeSharding or InitializeSharding | ||
| // from this constructor. Derived class state is not fully initialized at this point. | ||
| // Instead, derived classes should invoke InitializeAfterConstruction() at the end | ||
| // of their constructors once their own state is fully set up. | ||
| } | ||
|
|
||
| /// <summary> | ||
| /// Performs the sharding initialization sequence after the object is fully constructed. | ||
| /// </summary> | ||
| /// <remarks> | ||
| /// Derived classes should call this method at the end of their constructors, after | ||
| /// initializing any state that is required by OnBeforeInitializeSharding or | ||
| /// InitializeSharding. | ||
| /// </remarks> | ||
| protected void InitializeAfterConstruction() | ||
| { | ||
| // Allow derived classes to set up state before sharding | ||
| OnBeforeInitializeSharding(); | ||
| InitializeSharding(); |
Summary
Test Coverage
Bug Fixes
_numStageswas 0 whenInitializeSharding()was calledOnBeforeInitializeSharding()to set fields from Config_tensorParallelSizewas 0 whenInitializeSharding()was calledOnBeforeInitializeSharding()to set fields from ConfigOnBeforeInitializeSharding()Root Cause: C# constructor initialization order - derived class fields are uninitialized when the base constructor calls virtual methods.
Test Plan
Closes #654
🤖 Generated with Claude Code