FazBrowse GitHub Viewer | Trending |
URL:
| Home
Tools: [Download Repo ZIP]   [Original HTTPS Page]

Add SameDiff control flow execution framework (#10368) · deeplearning4j/deeplearning4j@95d5708 · GitHub

Commit 95d5708

Browse files
andauthored
Add SameDiff control flow execution framework (#10368)
New classes for frame-aware loop and conditional execution: - Frame analysis (FrameAnalyzer, FrameExecutionOptimizer) - Loop state tracking (LoopInfo, LoopState, LoopAnalysisHelpers) - Cross-frame references (CrossFrameReference, CrossLoopAnalysis) - Loop termination analysis and error reporting - Execution visualization (FrameVisualizer, SameDiffExecutionVisualizer) Updates to existing execution infrastructure: - AbstractSession, InferenceSession - ForwardExecutionDAG, ExecutionNode - FlatBuffersMapper, SDZSerializer 🤖 Generated with [Claude Code](https://claude.com/claude-code) Co-authored-by: Claude Opus 4.5 <noreply@anthropic.com>
1 parent 47757e7 commit 95d5708

44 files changed

Lines changed: 10810 additions & 990 deletions

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.
Lines changed: 46 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,46 @@
1+
/*
2+
* ******************************************************************************
3+
* *
4+
* *
5+
* * This program and the accompanying materials are made available under the
6+
* * terms of the Apache License, Version 2.0 which is available at
7+
* * https://www.apache.org/licenses/LICENSE-2.0.
8+
* *
9+
* * See the NOTICE file distributed with this work for additional
10+
* * information regarding copyright ownership.
11+
* * Unless required by applicable law or agreed to in writing, software
12+
* * distributed under the License is distributed on an "AS IS" BASIS, WITHOUT
13+
* * WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the
14+
* * License for the specific language governing permissions and limitations
15+
* * under the License.
16+
* *
17+
* * SPDX-License-Identifier: Apache-2.0
18+
* *****************************************************************************
19+
*/
20+
21+
package org.nd4j.autodiff.samediff;
22+
23+
/**
24+
* Cross-frame variable reference information
25+
*/
26+
public class CrossFrameReference {
27+
public String variableName;
28+
public String sourceFrame;
29+
public String targetFrame;
30+
public int sourceIteration;
31+
public int targetIteration;
32+
public String mediatingOperation;
33+
public CrossFrameReferenceType referenceType;
34+
35+
public CrossFrameReference() {
36+
// Default constructor
37+
}
38+
39+
public CrossFrameReference(String variableName, String sourceFrame, String targetFrame,
40+
CrossFrameReferenceType referenceType) {
41+
this.variableName = variableName;
42+
this.sourceFrame = sourceFrame;
43+
this.targetFrame = targetFrame;
44+
this.referenceType = referenceType;
45+
}
46+
}
Lines changed: 32 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,32 @@
1+
/*
2+
* ******************************************************************************
3+
* *
4+
* *
5+
* * This program and the accompanying materials are made available under the
6+
* * terms of the Apache License, Version 2.0 which is available at
7+
* * https://www.apache.org/licenses/LICENSE-2.0.
8+
* *
9+
* * See the NOTICE file distributed with this work for additional
10+
* * information regarding copyright ownership.
11+
* * Unless required by applicable law or agreed to in writing, software
12+
* * distributed under the License is distributed on an "AS IS" BASIS, WITHOUT
13+
* * WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the
14+
* * License for the specific language governing permissions and limitations
15+
* * under the License.
16+
* *
17+
* * SPDX-License-Identifier: Apache-2.0
18+
* *****************************************************************************
19+
*/
20+
21+
package org.nd4j.autodiff.samediff;
22+
23+
/**
24+
* Types of cross-frame variable references
25+
*/
26+
public enum CrossFrameReferenceType {
27+
DIRECT, // Direct variable reference
28+
ENTER, // Variable entering frame
29+
EXIT, // Variable exiting frame
30+
LOOP_CARRIED, // Variable carried across loop iterations
31+
CONDITIONAL // Variable from conditional branch
32+
}
Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,16 @@
1+
package org.nd4j.autodiff.samediff;
2+
3+
import lombok.Data;
4+
5+
import java.util.ArrayList;
6+
import java.util.HashMap;
7+
import java.util.List;
8+
import java.util.Map;
9+
10+
@Data
11+
public class CrossLoopAnalysis {
12+
private Map<TerminationType, Long> terminationTypeDistribution = new HashMap<>();
13+
private List<String> terminationCorrelations = new ArrayList<>();
14+
private List<String> systemWideIssues = new ArrayList<>();
15+
private Map<String, Integer> commonProblematicVariables = new HashMap<>();
16+
}
Lines changed: 250 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,250 @@
1+
/*
2+
* ******************************************************************************
3+
* *
4+
* *
5+
* * This program and the accompanying materials are made available under the
6+
* * terms of the Apache License, Version 2.0 which is available at
7+
* * https://www.apache.org/licenses/LICENSE-2.0.
8+
* *
9+
* * See the NOTICE file distributed with this work for additional
10+
* * information regarding copyright ownership.
11+
* * Unless required by applicable law or agreed to in writing, software
12+
* * distributed under the License is distributed on an "AS IS" BASIS, WITHOUT
13+
* * WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the
14+
* * License for the specific language governing permissions and limitations
15+
* * under the License.
16+
* *
17+
* * SPDX-License-Identifier: Apache-2.0
18+
* *****************************************************************************
19+
*/
20+
21+
package org.nd4j.autodiff.samediff;
22+
23+
import java.util.*;
24+
25+
/**
26+
* Utility class for analyzing frame structures and dependencies in DAG execution plans
27+
*/
28+
public class FrameAnalyzer {
29+
30+
/**
31+
* Analyze frame execution patterns and detect potential issues
32+
*/
33+
public static FrameAnalysisResult analyzeFrameExecution(DAGExecutionPlan plan) {
34+
FrameAnalysisResult result = new FrameAnalysisResult();
35+
36+
// Analyze frame depth and nesting
37+
analyzeFrameNesting(plan, result);
38+
39+
// Detect frame dependency cycles
40+
detectFrameCycles(plan, result);
41+
42+
// Analyze frame transition patterns
43+
analyzeTransitionPatterns(plan, result);
44+
45+
// Check for frame isolation issues
46+
checkFrameIsolation(plan, result);
47+
48+
return result;
49+
}
50+
51+
/**
52+
* Find the critical path through frames
53+
*/
54+
public static List<String> findFrameCriticalPath(DAGExecutionPlan plan) {
55+
Map<String, Integer> frameOperationCounts = new HashMap<>();
56+
57+
for (Map.Entry<String, List<String>> entry : plan.getFrameExecutionOrder().entrySet()) {
58+
frameOperationCounts.put(entry.getKey(), entry.getValue().size());
59+
}
60+
61+
return frameOperationCounts.entrySet().stream()
62+
.sorted(Map.Entry.<String, Integer>comparingByValue().reversed())
63+
.map(Map.Entry::getKey)
64+
.collect(ArrayList::new, (list, item) -> list.add(item), ArrayList::addAll);
65+
}
66+
67+
/**
68+
* Calculate frame execution efficiency metrics
69+
*/
70+
public static Map<String, Double> calculateFrameEfficiency(DAGExecutionPlan plan) {
71+
Map<String, Double> efficiency = new HashMap<>();
72+
73+
for (String frameName : plan.getFrameMetadata().keySet()) {
74+
List<String> frameOps = plan.getOperationsInFrame(frameName);
75+
Set<String> frameVars = plan.getVariablesInFrame(frameName);
76+
77+
if (!frameOps.isEmpty()) {
78+
double opsToVarsRatio = frameVars.isEmpty() ? 0.0 : (double) frameOps.size() / frameVars.size();
79+
efficiency.put(frameName, opsToVarsRatio);
80+
}
81+
}
82+
83+
return efficiency;
84+
}
85+
86+
/**
87+
* Find frames that could be parallelized
88+
*/
89+
public static Set<Set<String>> findParallelizableFrames(DAGExecutionPlan plan) {
90+
Set<Set<String>> parallelGroups = new HashSet<>();
91+
Map<String, Set<String>> frameDeps = plan.analyzeFrameDependencies();
92+
93+
// Find frames at the same depth with no dependencies between them
94+
Map<Integer, Set<String>> framesByDepth = new HashMap<>();
95+
for (Map.Entry<String, FrameMetadata> entry : plan.getFrameMetadata().entrySet()) {
96+
framesByDepth.computeIfAbsent(entry.getValue().depth, k -> new HashSet<>()).add(entry.getKey());
97+
}
98+
99+
for (Set<String> framesAtDepth : framesByDepth.values()) {
100+
if (framesAtDepth.size() > 1) {
101+
Set<String> parallelizable = new HashSet<>();
102+
for (String frame : framesAtDepth) {
103+
boolean canParallelize = true;
104+
for (String otherFrame : framesAtDepth) {
105+
if (!frame.equals(otherFrame)) {
106+
Set<String> deps = frameDeps.getOrDefault(frame, Collections.emptySet());
107+
if (deps.contains(otherFrame)) {
108+
canParallelize = false;
109+
break;
110+
}
111+
}
112+
}
113+
if (canParallelize) {
114+
parallelizable.add(frame);
115+
}
116+
}
117+
if (parallelizable.size() > 1) {
118+
parallelGroups.add(parallelizable);
119+
}
120+
}
121+
}
122+
123+
return parallelGroups;
124+
}
125+
126+
private static void analyzeFrameNesting(DAGExecutionPlan plan, FrameAnalysisResult result) {
127+
int maxDepth = 0;
128+
Map<Integer, Integer> depthCounts = new HashMap<>();
129+
130+
for (FrameMetadata meta : plan.getFrameMetadata().values()) {
131+
maxDepth = Math.max(maxDepth, meta.depth);
132+
depthCounts.merge(meta.depth, 1, Integer::sum);
133+
}
134+
135+
result.maxNestingDepth = maxDepth;
136+
result.frameCountByDepth = depthCounts;
137+
138+
// Flag deeply nested frames as potential issues
139+
if (maxDepth > 5) {
140+
result.warnings.add("Deep frame nesting detected (depth: " + maxDepth + "). Consider flattening.");
141+
}
142+
}
143+
144+
private static void detectFrameCycles(DAGExecutionPlan plan, FrameAnalysisResult result) {
145+
Map<String, Set<String>> frameDeps = plan.getFrameDependencies();
146+
Set<String> visited = new HashSet<>();
147+
Set<String> recursionStack = new HashSet<>();
148+
149+
for (String frame : plan.getFrameMetadata().keySet()) {
150+
if (!visited.contains(frame)) {
151+
if (hasCycleDFS(frame, frameDeps, visited, recursionStack, result.frameCycles)) {
152+
result.hasCycles = true;
153+
}
154+
}
155+
}
156+
}
157+
158+
private static boolean hasCycleDFS(String frame, Map<String, Set<String>> deps,
159+
Set<String> visited, Set<String> stack, List<String> cycles) {
160+
visited.add(frame);
161+
stack.add(frame);
162+
163+
Set<String> frameDeps = deps.getOrDefault(frame, Collections.emptySet());
164+
for (String dep : frameDeps) {
165+
if (!visited.contains(dep)) {
166+
if (hasCycleDFS(dep, deps, visited, stack, cycles)) {
167+
return true;
168+
}
169+
} else if (stack.contains(dep)) {
170+
cycles.add("Cycle detected: " + frame + " -> " + dep);
171+
return true;
172+
}
173+
}
174+
175+
stack.remove(frame);
176+
return false;
177+
}
178+
179+
private static void analyzeTransitionPatterns(DAGExecutionPlan plan, FrameAnalysisResult result) {
180+
Map<FrameTransition, Integer> transitionCounts = new HashMap<>();
181+
182+
for (FrameMetadata meta : plan.getFrameMetadata().values()) {
183+
for (Map.Entry<FrameTransition, Integer> entry : meta.transitionCounts.entrySet()) {
184+
transitionCounts.merge(entry.getKey(), entry.getValue(), Integer::sum);
185+
}
186+
}
187+
188+
result.transitionPatterns = transitionCounts;
189+
190+
// Analyze patterns for potential optimizations
191+
int enterExitRatio = transitionCounts.getOrDefault(FrameTransition.ENTER, 0) -
192+
transitionCounts.getOrDefault(FrameTransition.EXIT, 0);
193+
if (Math.abs(enterExitRatio) > 5) {
194+
result.warnings.add("Unbalanced ENTER/EXIT transitions (difference: " + enterExitRatio + ")");
195+
}
196+
}
197+
198+
private static void checkFrameIsolation(DAGExecutionPlan plan, FrameAnalysisResult result) {
199+
for (String frameName : plan.getFrameMetadata().keySet()) {
200+
Set<String> inputs = plan.getFrameInputVariables().getOrDefault(frameName, Collections.emptySet());
201+
Set<String> outputs = plan.getFrameOutputVariables().getOrDefault(frameName, Collections.emptySet());
202+
203+
if (inputs.isEmpty() && outputs.isEmpty()) {
204+
List<String> frameOps = plan.getOperationsInFrame(frameName);
205+
if (!frameOps.isEmpty()) {
206+
result.isolatedFrames.add(frameName);
207+
}
208+
}
209+
}
210+
211+
if (!result.isolatedFrames.isEmpty()) {
212+
result.warnings.add("Found " + result.isolatedFrames.size() + " isolated frames with no external I/O");
213+
}
214+
}
215+
216+
/**
217+
* Result of frame analysis
218+
*/
219+
public static class FrameAnalysisResult {
220+
public int maxNestingDepth;
221+
public Map<Integer, Integer> frameCountByDepth = new HashMap<>();
222+
public boolean hasCycles = false;
223+
public List<String> frameCycles = new ArrayList<>();
224+
public Map<FrameTransition, Integer> transitionPatterns = new HashMap<>();
225+
public List<String> isolatedFrames = new ArrayList<>();
226+
public List<String> warnings = new ArrayList<>();
227+
228+
public boolean hasIssues() {
229+
return hasCycles || !isolatedFrames.isEmpty() || !warnings.isEmpty();
230+
}
231+
232+
public String getSummary() {
233+
StringBuilder sb = new StringBuilder();
234+
sb.append("Frame Analysis Summary:\n");
235+
sb.append(" Max nesting depth: ").append(maxNestingDepth).append("\n");
236+
sb.append(" Has cycles: ").append(hasCycles).append("\n");
237+
sb.append(" Isolated frames: ").append(isolatedFrames.size()).append("\n");
238+
sb.append(" Warnings: ").append(warnings.size()).append("\n");
239+
240+
if (!warnings.isEmpty()) {
241+
sb.append("\nWarnings:\n");
242+
for (String warning : warnings) {
243+
sb.append(" - ").append(warning).append("\n");
244+
}
245+
}
246+
247+
return sb.toString();
248+
}
249+
}
250+
}
Lines changed: 18 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,18 @@
1+
package org.nd4j.autodiff.samediff;
2+
3+
import lombok.Data;
4+
5+
import java.util.ArrayList;
6+
import java.util.HashMap;
7+
import java.util.List;
8+
import java.util.Map;
9+
10+
@Data
11+
public class FrameContextInfo {
12+
private String frameName;
13+
private int iteration;
14+
private String parentFrame;
15+
private int nestingDepth;
16+
private List<String> relatedFrames = new ArrayList<>();
17+
private Map<String, List<String>> crossFrameReferences = new HashMap<>();
18+
}

0 commit comments

Comments
 (0)

Back | FazBrowse Home | New Git URL