Generates a computation graph starting from the variable itself. :param print_gradients: A boolean indicating whether to print gradients in the graph. :return: A visualization of the computation graph.
(self, print_gradients: bool=False)
| 189 | v.grad_fn(backward_engine=backward_engine) |
| 190 | |
| 191 | def generate_graph(self, print_gradients: bool=False): |
| 192 | """ |
| 193 | Generates a computation graph starting from the variable itself. |
| 194 | |
| 195 | :param print_gradients: A boolean indicating whether to print gradients in the graph. |
| 196 | :return: A visualization of the computation graph. |
| 197 | """ |
| 198 | try: |
| 199 | from graphviz import Digraph |
| 200 | except ImportError: |
| 201 | raise ImportError("Please install graphviz to visualize the computation graphs. You can install it using `pip install graphviz`.") |
| 202 | |
| 203 | def wrap_text(text, width=40): |
| 204 | """Wraps text at a given number of characters using HTML line breaks.""" |
| 205 | words = text.split() |
| 206 | wrapped_text = "" |
| 207 | line = "" |
| 208 | for word in words: |
| 209 | if len(line) + len(word) + 1 > width: |
| 210 | wrapped_text += line + "<br/>" |
| 211 | line = word |
| 212 | else: |
| 213 | if line: |
| 214 | line += " " |
| 215 | line += word |
| 216 | wrapped_text += line |
| 217 | return wrapped_text |
| 218 | |
| 219 | def wrap_and_escape(text, width=40): |
| 220 | return wrap_text(text.replace("<", "<").replace(">", ">"), width) |
| 221 | |
| 222 | topo = [] |
| 223 | visited = set() |
| 224 | def build_topo(v): |
| 225 | if v not in visited: |
| 226 | visited.add(v) |
| 227 | for predecessor in v.predecessors: |
| 228 | build_topo(predecessor) |
| 229 | topo.append(v) |
| 230 | |
| 231 | def get_grad_fn_name(name): |
| 232 | ws = name.split(" ") |
| 233 | ws = [w for w in ws if "backward" in w] |
| 234 | return " ".join(ws) |
| 235 | |
| 236 | build_topo(self) |
| 237 | |
| 238 | graph = Digraph(comment='Computation Graph starting from {}'.format(self.role_description)) |
| 239 | graph.attr(rankdir='TB') # Set the graph direction from top to bottom |
| 240 | graph.attr(ranksep='0.2') # Adjust the spacing between ranks |
| 241 | graph.attr(bgcolor='lightgrey') # Set the background color of the graph |
| 242 | graph.attr(fontsize='7.5') # Set the font size of the graph |
| 243 | |
| 244 | for v in reversed(topo): |
| 245 | # Add each node to the graph |
| 246 | label_color = 'darkblue' |
| 247 | |
| 248 | node_label = ( |