| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 47757e7 commit 95d5708
44 files changed
| Original file line number | Diff line number | Diff 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 | + } | ||
| Original file line number | Diff line number | Diff 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 | + } | ||
| Original file line number | Diff line number | Diff 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 | + } | ||
| Original file line number | Diff line number | Diff 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 | + } | ||
| Original file line number | Diff line number | Diff 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 | + } | ||
| Back | FazBrowse Home | New Git URL |
0 commit comments