Skip to content

Commit 7058bda

Browse files
committed
formatted with ruff
1 parent 4b83219 commit 7058bda

4 files changed

Lines changed: 132 additions & 78 deletions

File tree

examples/graphs.ipynb

Lines changed: 121 additions & 71 deletions
Original file line numberDiff line numberDiff line change
@@ -16,51 +16,55 @@
1616
"metadata": {},
1717
"outputs": [],
1818
"source": [
19-
"opt_weighted = np.array([\n",
20-
" 0.833785879763828,\n",
21-
" 0.8474403585597967,\n",
22-
" 0.9333143427873113,\n",
23-
" 0.8901245335864413,\n",
24-
" 0.8377901148644222,\n",
25-
" 0.8478829150746583,\n",
26-
" 0.8165234517139993,\n",
27-
" 0.8499555079041013,\n",
28-
" 0.835089689986721,\n",
29-
" 0.8381280690241771,\n",
30-
" 0.8672643015339117,\n",
31-
" 0.8785954542980072,\n",
32-
" 0.819237021124287,\n",
33-
" 0.8371113686804559,\n",
34-
" 0.863054297911893,\n",
35-
" 0.8227415494858079,\n",
36-
" 0.8372289303002084,\n",
37-
" 0.8404439216966082,\n",
38-
" 0.8276779336534488,\n",
39-
" 0.8799974291947237\n",
40-
"])\n",
19+
"opt_weighted = np.array(\n",
20+
" [\n",
21+
" 0.833785879763828,\n",
22+
" 0.8474403585597967,\n",
23+
" 0.9333143427873113,\n",
24+
" 0.8901245335864413,\n",
25+
" 0.8377901148644222,\n",
26+
" 0.8478829150746583,\n",
27+
" 0.8165234517139993,\n",
28+
" 0.8499555079041013,\n",
29+
" 0.835089689986721,\n",
30+
" 0.8381280690241771,\n",
31+
" 0.8672643015339117,\n",
32+
" 0.8785954542980072,\n",
33+
" 0.819237021124287,\n",
34+
" 0.8371113686804559,\n",
35+
" 0.863054297911893,\n",
36+
" 0.8227415494858079,\n",
37+
" 0.8372289303002084,\n",
38+
" 0.8404439216966082,\n",
39+
" 0.8276779336534488,\n",
40+
" 0.8799974291947237,\n",
41+
" ]\n",
42+
")\n",
4143
"\n",
42-
"opt_unweighted = np.array([\n",
43-
" 0.8646741695464797,\n",
44-
" 0.8063692275885563,\n",
45-
" 0.8275313927869129,\n",
46-
" 0.8697326449072285,\n",
47-
" 0.8201401404514435,\n",
48-
" 0.8163135245347478,\n",
49-
" 0.8521422155343803,\n",
50-
" 0.8544577303206086,\n",
51-
" 0.854259427565678,\n",
52-
" 0.8650212587824293,\n",
53-
" 0.8328442228068212,\n",
54-
" 0.8063805930933375,\n",
55-
" 0.8221443219549337,\n",
56-
" 0.8779375010235294,\n",
57-
" 0.8375775878596458,\n",
58-
" 0.8453875719362004,\n",
59-
" 0.9377081010751663,\n",
60-
" 0.8246441791012029,\n",
61-
" 0.8780643264199518,\n",
62-
" 0.840796954692549\n",
63-
"])"
44+
"opt_unweighted = np.array(\n",
45+
" [\n",
46+
" 0.8646741695464797,\n",
47+
" 0.8063692275885563,\n",
48+
" 0.8275313927869129,\n",
49+
" 0.8697326449072285,\n",
50+
" 0.8201401404514435,\n",
51+
" 0.8163135245347478,\n",
52+
" 0.8521422155343803,\n",
53+
" 0.8544577303206086,\n",
54+
" 0.854259427565678,\n",
55+
" 0.8650212587824293,\n",
56+
" 0.8328442228068212,\n",
57+
" 0.8063805930933375,\n",
58+
" 0.8221443219549337,\n",
59+
" 0.8779375010235294,\n",
60+
" 0.8375775878596458,\n",
61+
" 0.8453875719362004,\n",
62+
" 0.9377081010751663,\n",
63+
" 0.8246441791012029,\n",
64+
" 0.8780643264199518,\n",
65+
" 0.840796954692549,\n",
66+
" ]\n",
67+
")"
6468
]
6569
},
6670
{
@@ -99,13 +103,13 @@
99103
" ol1.append(min(ol1[i - 1], opt_weighted[i]))\n",
100104
" ol2.append(min(ol2[i - 1], opt_unweighted[i]))\n",
101105
"\n",
102-
"ax1.plot(np.arange(len(ol1)), ol1, '-o')\n",
103-
"ax1.plot(np.arange(len(ol2)), ol2, '-o')\n",
104-
"ax1.set_xlabel('# of iterations')\n",
105-
"ax1.set_ylabel('Testing Loss')\n",
106+
"ax1.plot(np.arange(len(ol1)), ol1, \"-o\")\n",
107+
"ax1.plot(np.arange(len(ol2)), ol2, \"-o\")\n",
108+
"ax1.set_xlabel(\"# of iterations\")\n",
109+
"ax1.set_ylabel(\"Testing Loss\")\n",
106110
"ax1.legend([\"Weighted\", \"Unweighted\"])\n",
107111
"ax1.set_xticks([0, 5, 10, 15, 20])\n",
108-
"ax1.set_title('')"
112+
"ax1.set_title(\"\")"
109113
]
110114
},
111115
{
@@ -137,13 +141,13 @@
137141
"source": [
138142
"fig2, ax2 = plt.subplots()\n",
139143
"\n",
140-
"ax2.plot(np.arange(len(opt_unweighted)), opt_unweighted, '-o')\n",
141-
"ax2.plot(np.arange(len(opt_weighted)), opt_weighted, '-o')\n",
142-
"ax2.set_xlabel('# of iterations')\n",
143-
"ax2.set_ylabel('Testing Loss')\n",
144+
"ax2.plot(np.arange(len(opt_unweighted)), opt_unweighted, \"-o\")\n",
145+
"ax2.plot(np.arange(len(opt_weighted)), opt_weighted, \"-o\")\n",
146+
"ax2.set_xlabel(\"# of iterations\")\n",
147+
"ax2.set_ylabel(\"Testing Loss\")\n",
144148
"ax2.legend([\"Weighted\", \"Unweighted\"])\n",
145149
"ax2.set_xticks([0, 5, 10, 15, 20])\n",
146-
"ax2.set_title('')"
150+
"ax2.set_title(\"\")"
147151
]
148152
},
149153
{
@@ -152,9 +156,55 @@
152156
"metadata": {},
153157
"outputs": [],
154158
"source": [
155-
"bk_weighted = np.array([0.8711097403696388, 0.8219735683149593, 0.9250104393169378, 0.928701077297235, 0.8662791377419878, 0.8644050124344552, 0.8295078087764182, 0.8913433498637692, 0.7863149593590172, 0.8447837476517744, 0.8883395092502521, 0.8882391870401467, 0.8522632577616698, 0.884719618946124, 0.876211231681192, 0.8631530969765535, 0.8744676268784104, 0.7886375295128792, 0.8420131016688742, 0.8851713993746764])\n",
159+
"bk_weighted = np.array(\n",
160+
" [\n",
161+
" 0.8711097403696388,\n",
162+
" 0.8219735683149593,\n",
163+
" 0.9250104393169378,\n",
164+
" 0.928701077297235,\n",
165+
" 0.8662791377419878,\n",
166+
" 0.8644050124344552,\n",
167+
" 0.8295078087764182,\n",
168+
" 0.8913433498637692,\n",
169+
" 0.7863149593590172,\n",
170+
" 0.8447837476517744,\n",
171+
" 0.8883395092502521,\n",
172+
" 0.8882391870401467,\n",
173+
" 0.8522632577616698,\n",
174+
" 0.884719618946124,\n",
175+
" 0.876211231681192,\n",
176+
" 0.8631530969765535,\n",
177+
" 0.8744676268784104,\n",
178+
" 0.7886375295128792,\n",
179+
" 0.8420131016688742,\n",
180+
" 0.8851713993746764,\n",
181+
" ]\n",
182+
")\n",
156183
"\n",
157-
"bk_unweighted = np.array([0.8768957655900603, 0.8436034621706434, 0.8426990888680622, 1.0322862487689706, 0.8425512856738583, 0.9015028643759952, 0.849575503616576, 0.8581730317158304, 1.0051764811679815, 0.8984316902555478, 0.7911360404294008, 0.8346848123392482, 0.8960412535697792, 0.891405322749144, 0.8127066090608098, 1.0492611349008645, 0.8460451801111744, 0.9509688987853421, 0.8639677797153498, 0.8604594484256332])"
184+
"bk_unweighted = np.array(\n",
185+
" [\n",
186+
" 0.8768957655900603,\n",
187+
" 0.8436034621706434,\n",
188+
" 0.8426990888680622,\n",
189+
" 1.0322862487689706,\n",
190+
" 0.8425512856738583,\n",
191+
" 0.9015028643759952,\n",
192+
" 0.849575503616576,\n",
193+
" 0.8581730317158304,\n",
194+
" 1.0051764811679815,\n",
195+
" 0.8984316902555478,\n",
196+
" 0.7911360404294008,\n",
197+
" 0.8346848123392482,\n",
198+
" 0.8960412535697792,\n",
199+
" 0.891405322749144,\n",
200+
" 0.8127066090608098,\n",
201+
" 1.0492611349008645,\n",
202+
" 0.8460451801111744,\n",
203+
" 0.9509688987853421,\n",
204+
" 0.8639677797153498,\n",
205+
" 0.8604594484256332,\n",
206+
" ]\n",
207+
")"
158208
]
159209
},
160210
{
@@ -193,14 +243,14 @@
193243
" bk1.append(min(bk1[i - 1], bk_weighted[i]))\n",
194244
" bk2.append(min(bk2[i - 1], bk_unweighted[i]))\n",
195245
"\n",
196-
"ax3.plot(np.arange(len(bk1)), bk1, '-o')\n",
246+
"ax3.plot(np.arange(len(bk1)), bk1, \"-o\")\n",
197247
"# ax3.plot(np.arange(len(ol1)), ol1, '-o')\n",
198-
"ax3.plot(np.arange(len(bk2)), bk2, '-o')\n",
199-
"ax3.set_xlabel('# of iterations')\n",
200-
"ax3.set_ylabel('Testing Loss')\n",
248+
"ax3.plot(np.arange(len(bk2)), bk2, \"-o\")\n",
249+
"ax3.set_xlabel(\"# of iterations\")\n",
250+
"ax3.set_ylabel(\"Testing Loss\")\n",
201251
"ax3.legend([\"Weighted\", \"Unweighted\"])\n",
202252
"ax3.set_xticks([0, 5, 10, 15, 20])\n",
203-
"ax3.set_title('')\n"
253+
"ax3.set_title(\"\")"
204254
]
205255
},
206256
{
@@ -232,13 +282,13 @@
232282
"source": [
233283
"fig4, ax4 = plt.subplots()\n",
234284
"\n",
235-
"ax4.plot(np.arange(len(bk_unweighted)), bk_unweighted, '-o')\n",
236-
"ax4.plot(np.arange(len(bk_weighted)), bk_weighted, '-o')\n",
237-
"ax4.set_xlabel('# of iterations')\n",
238-
"ax4.set_ylabel('Testing Loss')\n",
285+
"ax4.plot(np.arange(len(bk_unweighted)), bk_unweighted, \"-o\")\n",
286+
"ax4.plot(np.arange(len(bk_weighted)), bk_weighted, \"-o\")\n",
287+
"ax4.set_xlabel(\"# of iterations\")\n",
288+
"ax4.set_ylabel(\"Testing Loss\")\n",
239289
"ax4.legend([\"Weighted\", \"Unweighted\"])\n",
240290
"ax4.set_xticks([0, 5, 10, 15, 20])\n",
241-
"ax4.set_title('')"
291+
"ax4.set_title(\"\")"
242292
]
243293
},
244294
{
@@ -271,13 +321,13 @@
271321
"fig5, ax5 = plt.subplots()\n",
272322
"\n",
273323
"\n",
274-
"ax5.plot(np.arange(len(bk1)), bk1, '-o')\n",
275-
"ax5.plot(np.arange(len(ol1)), ol1, '-o')\n",
276-
"ax5.set_xlabel('# of iterations')\n",
277-
"ax3.set_ylabel('Testing Loss')\n",
324+
"ax5.plot(np.arange(len(bk1)), bk1, \"-o\")\n",
325+
"ax5.plot(np.arange(len(ol1)), ol1, \"-o\")\n",
326+
"ax5.set_xlabel(\"# of iterations\")\n",
327+
"ax3.set_ylabel(\"Testing Loss\")\n",
278328
"ax5.legend([\"Backtracking Weighted\", \"Naive Weighted\"])\n",
279329
"ax5.set_xticks([0, 5, 10, 15, 20])\n",
280-
"ax5.set_title('')\n"
330+
"ax5.set_title(\"\")"
281331
]
282332
},
283333
{

examples/searchspace.ipynb

Lines changed: 9 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,7 @@
1818
"source": [
1919
"distances, idxs = construct_graph_search_space(10)\n",
2020
"graph = nx.from_numpy_array(distances)\n",
21-
"colors = [get_color(graph.get_edge_data(*edge)['weight']) for edge in graph.edges()]"
21+
"colors = [get_color(graph.get_edge_data(*edge)[\"weight\"]) for edge in graph.edges()]"
2222
]
2323
},
2424
{
@@ -38,7 +38,7 @@
3838
}
3939
],
4040
"source": [
41-
"nx.draw(graph, node_size=10, node_color='black', edge_color=colors)"
41+
"nx.draw(graph, node_size=10, node_color=\"black\", edge_color=colors)"
4242
]
4343
},
4444
{
@@ -58,7 +58,13 @@
5858
}
5959
],
6060
"source": [
61-
"nx.draw(graph, pos=nx.spring_layout(graph), node_size=10, node_color='black', edge_color=colors)"
61+
"nx.draw(\n",
62+
" graph,\n",
63+
" pos=nx.spring_layout(graph),\n",
64+
" node_size=10,\n",
65+
" node_color=\"black\",\n",
66+
" edge_color=colors,\n",
67+
")"
6268
]
6369
},
6470
{

nlgm/manifolds.py

Lines changed: 1 addition & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -54,9 +54,7 @@ def exponential_map(self, tangent_vector: torch.Tensor) -> torch.Tensor:
5454
"""
5555
pass
5656

57-
def distance(
58-
self, point_x: torch.Tensor, point_y: torch.Tensor
59-
) -> torch.Tensor:
57+
def distance(self, point_x: torch.Tensor, point_y: torch.Tensor) -> torch.Tensor:
6058
"""
6159
Compute geodesic distance between two points on the manifold.
6260

nlgm/train.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -69,7 +69,7 @@ def train_and_evaluate(
6969
test_loss /= len(test_loader)
7070

7171
print(
72-
f"Epoch [{epoch+1}/{epochs}], Train Loss: {train_loss:.4f}, Test Loss: {test_loss:.4f}"
72+
f"Epoch [{epoch + 1}/{epochs}], Train Loss: {train_loss:.4f}, Test Loss: {test_loss:.4f}"
7373
)
7474

7575
return train_losses, test_loss

0 commit comments

Comments
 (0)