@@ -1041,30 +1041,38 @@ IRType VisitD2D(DistributedType inv, DistributedType outv)
10411041 }
10421042
10431043 var partialDims = new List < int > ( ) ;
1044- for ( int i = 0 ; i < inv . AxisPolicies . Count ; i ++ )
1044+ if ( inv . Partial is not null )
10451045 {
1046- if ( inv . Partial is not null && outv . AxisPolicies [ i ] is SBPSplit s )
1046+ for ( int i = 0 ; i < inv . AxisPolicies . Count ; i ++ )
10471047 {
1048- if ( inv . AxisPolicies [ i ] is SBPSplit splitIn )
1048+ if ( inv . AxisPolicies [ i ] is SBPSplit && outv . AxisPolicies [ i ] is SBPBroadCast )
10491049 {
1050- if ( splitIn . Axes . Except ( s . Axes ) . Any ( ) )
1051- {
1052- return new InvalidType ( "Not Supported Split-> Split." ) ;
1053- }
1050+ return new InvalidType ( "Not supported input is BroadCast output is Split" ) ;
10541051 }
10551052
1056- if ( s . Axes . Except ( inv . Partial . Axes ) . ToArray ( ) != s . Axes )
1053+ if ( outv . AxisPolicies [ i ] is SBPSplit s )
10571054 {
1058- if ( s . Axes . Except ( inv . Partial . Axes ) . Any ( ) )
1055+ if ( inv . AxisPolicies [ i ] is SBPSplit splitIn )
10591056 {
1060- return new InvalidType ( "Not Supported Partial-> Split." ) ;
1057+ if ( splitIn . Axes . Except ( s . Axes ) . Any ( ) )
1058+ {
1059+ return new InvalidType ( "Not Supported Split-> Split." ) ;
1060+ }
10611061 }
1062- else
1062+
1063+ if ( s . Axes . Except ( inv . Partial . Axes ) . ToArray ( ) != s . Axes )
10631064 {
10641065 partialDims . Add ( i ) ;
10651066 }
10661067 }
10671068 }
1069+
1070+ var ndspsIn = DistributedUtility . AxisPolicesToNDSBP ( inv . AxisPolicies , inv . Placement . Rank ) ;
1071+ var ndspsOut = DistributedUtility . AxisPolicesToNDSBP ( outv . AxisPolicies , outv . Placement . Rank ) ;
1072+ if ( Enumerable . Range ( 0 , ndspsIn . Count ) . Any ( i => ndspsIn [ i ] is SBPSplit si && ( ndspsOut [ i ] is SBPBroadCast || ( ndspsOut [ i ] is SBPSplit so && so . Axes [ 0 ] != si . Axes [ 0 ] ) ) ) )
1073+ {
1074+ return new InvalidType ( "Not Supported Split-> Broadcast." ) ;
1075+ }
10681076 }
10691077
10701078 if ( partialDims . Count > 0 && ! Enumerable . Range ( 0 , inv . AxisPolicies . Count ) . Except ( partialDims . ToArray ( ) ) . All ( i => DistributedUtility . IsSamePolicy ( inv . AxisPolicies [ i ] , outv . AxisPolicies [ i ] ) ) )
0 commit comments