-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtest.py
More file actions
267 lines (237 loc) · 12 KB
/
Copy pathtest.py
File metadata and controls
267 lines (237 loc) · 12 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
import numpy as np
import pandas as pd
from sklearn.decomposition import PCA
from sklearn.preprocessing import StandardScaler
import matplotlib.pyplot as plt
import seaborn as sns
from scipy.spatial.distance import cdist
from sklearn.cluster import KMeans
import os
# 创建img文件夹(如果不存在)
os.makedirs('./img', exist_ok=True)
# 1. 手动定义25个任务的embedding数据
task_embeds = np.array([
[-0.20658779, 0.2779107 , 0.44654372, 0.11531506, -0.08924519, -0.22488531,
-0.38851383, 0.15576321, 0.34279835, -0.12856865, -0.5096578 , 0.00859224,
-0.06998068, 0.11024848, 0.12372816, -0.08613231],
[-0.26100153, 0.19054267, 0.24659978, 0.12588952, -0.04061749, -0.15205392,
-0.2968075 , 0.12696552, 0.22933045, -0.39477378, -0.622989 , 0.12019093,
0.02722398, 0.09794052, -0.06741637, -0.25307822],
[ 0.0277089 , 0.16985948, 0.31010774, 0.3941706 , 0.24748845, -0.24651477,
-0.08486969, 0.2475317 , 0.518292 , -0.04558394, -0.17490233, -0.25032434,
0.01976901, 0.38828373, 0.10160688, -0.05395523],
[-0.24007979, 0.29678413, 0.33310795, 0.2652036 , 0.19657671, -0.0252249 ,
-0.36024806, 0.14867222, 0.39272135, -0.08113212, -0.43978542, 0.15385497,
-0.09015198, 0.23796408, 0.19754404, 0.01238327],
[-0.04426052, 0.18152548, 0.34622523, 0.07696801, -0.05191238, -0.03663837,
-0.35776147, 0.26997316, 0.40244225, 0.02061555, -0.56725293, 0.17924197,
0.04373837, 0.25345716, 0.22786804, -0.00341963],
[-0.27413374, 0.20408747, 0.30842337, 0.02749732, -0.12465985, -0.05629698,
-0.45857552, 0.25191647, 0.32476622, -0.24290471, -0.46525502, 0.06331718,
0.19432656, 0.19877116, 0.16665483, -0.06972348],
[-0.37933806, 0.2051156 , 0.21455353, 0.15214378, 0.13351521, -0.316952 ,
-0.3581948 , 0.15057042, 0.2755774 , -0.43606362, -0.34331542, 0.14133869,
0.07164317, -0.003079 , 0.25168854, -0.05624598],
[-0.26686135, 0.30703661, 0.35100594, 0.16672018, -0.10206027, -0.29652926,
-0.37718943, 0.09700765, 0.26816997, -0.21433562, -0.4949916 , -0.00040094,
-0.03872437, 0.03640487, 0.1872678 , -0.18088309],
[-0.14707716, 0.07524768, 0.27254996, 0.08029824, -0.34369567, -0.05760436,
-0.46460047, 0.07691892, 0.36679804, -0.18110721, -0.5458974 , 0.07810631,
0.14458823, 0.1684058 , 0.16567585, 0.02448798],
[-0.27825493, 0.20392978, 0.42831054, 0.06814642, 0.08157098, -0.02314332,
-0.46722075, 0.20152666, 0.34225655, -0.20362069, -0.41530913, -0.09050969,
0.14504416, 0.21885107, 0.13606903, -0.00850486],
[-0.4110415 , 0.26482543, 0.24697165, 0.15508224, 0.07900012, -0.22508042,
-0.34200096, 0.0550651 , 0.15171692, -0.48700094, -0.43664575, 0.02089557,
-0.05746901, 0.13706435, 0.14024544, -0.07697841],
[-0.151193 , 0.26203293, 0.430212 , 0.1414277 , 0.04035709, -0.11919992,
-0.29544076, 0.3351954 , 0.46598494, -0.05519593, -0.4023197 , 0.12817417,
0.0321364 , 0.17019552, 0.18318896, -0.16087584],
[-0.29310235, 0.12902437, -0.10073388, -0.10336681, -0.09789649, -0.5123325 ,
-0.4311748 , -0.11158921, 0.41574585, -0.37439358, 0.12768319, -0.13196795,
-0.14106587, -0.02944577, 0.13772266, -0.14039882],
[-0.01062611, 0.21098052, 0.34865892, 0.35707593, 0.19615626, -0.2716436 ,
-0.14848457, 0.3261523 , 0.48390764, -0.0353336 , -0.18945248, -0.17342035,
-0.074489 , 0.3852715 , 0.10053387, -0.01207652],
[ 0.01882944, 0.44416195, 0.51997745, -0.04411813, -0.03236341, 0.04299548,
-0.32957393, 0.15747115, 0.13609882, -0.06529382, -0.2232569 , 0.06416909,
0.20101379, 0.371575 , -0.09426389, 0.36003184],
[-0.22127299, 0.04680366, 0.09543736, 0.05726872, 0.52671903, -0.12245892,
-0.05758868, 0.11568617, 0.3666006 , 0.03184023, 0.19949096, -0.3674161 ,
0.08282766, 0.4791076 , 0.07774743, 0.27327737],
[-0.29374915, 0.09001489, 0.18333401, 0.24202512, -0.19571564, 0.03808359,
-0.37917897, 0.06398076, 0.30603817, 0.17623907, -0.5227185 , 0.05619693,
0.01534921, -0.11835466, 0.4586562 , 0.00865489],
[-0.11391181, 0.18899922, 0.396368 , 0.11990426, -0.04686523, 0.06457673,
-0.3754157 , 0.2167931 , 0.4064403 , 0.10914782, -0.5423552 , 0.13920853,
0.02248408, 0.25906473, 0.16508032, 0.00309534],
[-0.18094493, 0.20871755, 0.3065863 , 0.02823077, -0.04271932, -0.17292623,
-0.29212803, 0.25596258, 0.47306618, -0.1756886 , -0.44039583, 0.20151499,
0.09081338, 0.07205759, 0.33735266, -0.17279679],
[-0.47010055, 0.14821944, 0.05413824, 0.23291151, 0.07054626, -0.08195179,
-0.16563812, 0.04469781, 0.36813262, -0.21449876, -0.5834188 , 0.18801498,
-0.07439016, 0.2462411 , 0.16340601, -0.09274971],
[-0.3719456 , 0.16554134, 0.18510932, 0.15909266, -0.10207097, -0.19506988,
-0.30109644, 0.05061112, 0.19228189, -0.43406194, -0.60012054, -0.11793538,
-0.04037042, 0.13608934, 0.04043341, -0.10846651],
[-0.20005916, 0.12192281, 0.46184894, 0.22291037, 0.1841017 , 0.10880456,
-0.3044917 , 0.17148411, 0.3578685 , 0.16291589, -0.5582695 , 0.04047256,
-0.10817407, 0.18282537, 0.0186245 , 0.02959633],
[-0.25887462, 0.14405802, 0.21609996, -0.05210091, -0.1019304 , -0.15784346,
-0.32028073, 0.16056903, 0.40054825, -0.34990394, -0.497256 , 0.07012419,
0.19760694, 0.23202 , 0.22138958, -0.14900781],
[-0.3956001 , 0.18216306, 0.23540016, 0.12328105, 0.09189074, -0.28621504,
-0.39284328, 0.18539518, 0.2527599 , -0.42407206, -0.37712777, 0.11969062,
0.10692398, 0.00756225, 0.215035 , -0.05137618],
[-0.15092975, 0.19305554, 0.22423108, 0.06672397, -0.23483588, -0.14714284,
-0.45997998, 0.22691585, 0.31760728, -0.25345474, -0.4799498 , 0.11686258,
0.2534718 , 0.10089111, 0.21940342, -0.11687464]
])
# 2. 定义任务名称(MT25_V3)
task_names = [
"reach-v3", "push-v3", "pick-place-v3", "door-open-v3",
"drawer-open-v3", "drawer-close-v3", "button-press-topdown-v3",
"peg-insert-side-v3", "window-open-v3", "window-close-v3",
"coffee-pull-v3", "pick-out-of-hole-v3", "disassemble-v3",
"pick-place-wall-v3", "basketball-v3", "stick-pull-v3",
"button-press-wall-v3", "faucet-open-v3", "door-lock-v3",
"lever-pull-v3", "sweep-into-v3", "faucet-close-v3",
"coffee-button-v3", "button-press-topdown-wall-v3", "dial-turn-v3",
]
# 3. 数据标准化
# scaler = StandardScaler()
# embeds_scaled = scaler.fit_transform(task_embeds)
embeds_scaled = task_embeds
# 4. 执行PCA
pca = PCA(n_components=2)
embeds_2d = pca.fit_transform(embeds_scaled)
# 5. 创建结果DataFrame
df = pd.DataFrame({
'Task': task_names,
'PC1': embeds_2d[:, 0],
'PC2': embeds_2d[:, 1]
})
# 6. 可视化1: 基础PCA散点图
plt.figure(figsize=(16, 12))
# 绘制散点图
scatter = plt.scatter(df['PC1'], df['PC2'],
c=range(len(task_names)),
cmap='tab20',
s=150,
alpha=0.7,
edgecolors='black',
linewidth=0.5)
# 添加任务标签(优化位置避免重叠)
for i, task in enumerate(task_names):
# 获取任务简称(去掉-v3后缀)
short_name = task.replace('-v3', '')
# 根据位置调整标注方向
offset_x = 6 if embeds_2d[i, 0] > 0 else -60
offset_y = 6 if embeds_2d[i, 1] > 0 else -10
plt.annotate(short_name,
(df['PC1'][i], df['PC2'][i]),
xytext=(offset_x, offset_y),
textcoords='offset points',
fontsize=8,
alpha=0.9,
fontweight='bold',
bbox=dict(boxstyle='round,pad=0.2', facecolor='white', alpha=0.7))
plt.xlabel(f'PC1 ({pca.explained_variance_ratio_[0]:.1%} variance)', fontsize=12)
plt.ylabel(f'PC2 ({pca.explained_variance_ratio_[1]:.1%} variance)', fontsize=12)
plt.title('MT25 Task Embeddings PCA Analysis\nTask Similarity Visualization',
fontsize=14, fontweight='bold')
plt.grid(True, alpha=0.3)
cbar = plt.colorbar(scatter, label='Task Index', shrink=0.8)
cbar.ax.tick_params(labelsize=10)
plt.tight_layout()
plt.savefig('./img/mt25_pca_base.png', dpi=300, bbox_inches='tight')
plt.close()
print("✓ 基础PCA图已保存到: ./img/mt25_pca_base.png")
# 7. 计算任务相似度矩阵
similarity_matrix = 1 - cdist(task_embeds, task_embeds, metric='cosine')
# 8. 找出最相似的任务对
def find_most_similar_tasks(sim_matrix, task_names, top_k=10):
"""找出最相似的任务对"""
similar_pairs = []
for i in range(len(task_names)):
for j in range(i+1, len(task_names)):
similar_pairs.append((task_names[i], task_names[j], sim_matrix[i][j]))
similar_pairs.sort(key=lambda x: x[2], reverse=True)
return similar_pairs[:top_k]
print("\n" + "="*70)
print("最相似的任务对(Top 10)")
print("="*70)
top_pairs = find_most_similar_tasks(similarity_matrix, task_names)
for task1, task2, similarity in top_pairs:
print(f"{task1:<30} ↔ {task2:<30} : 相似度 = {similarity:.4f}")
# 9. 输出每个任务在PCA空间中的坐标
print("\n" + "="*70)
print("PCA坐标详情")
print("="*70)
df_display = df.copy()
df_display['Task'] = df_display['Task'].str.replace('-v3', '')
print(df_display.round(4).to_string(index=False))
# 10. 任务聚类分析
n_clusters = 8
kmeans = KMeans(n_clusters=n_clusters, random_state=42)
clusters = kmeans.fit_predict(task_embeds)
df['Cluster'] = clusters
print("\n" + "="*70)
print("任务聚类结果(8个群组)")
print("="*70)
for cluster_id in range(n_clusters):
tasks_in_cluster = df[df['Cluster'] == cluster_id]['Task'].tolist()
task_names_clean = [t.replace('-v3', '') for t in tasks_in_cluster]
print(f"\nCluster {cluster_id}: {', '.join(task_names_clean)}")
# 11. 可视化2: 聚类结果图
plt.figure(figsize=(16, 12))
# 为每个聚类使用不同颜色
colors = plt.cm.Accent(np.linspace(0, 1, n_clusters))
for cluster_id in range(n_clusters):
cluster_data = df[df['Cluster'] == cluster_id]
plt.scatter(cluster_data['PC1'], cluster_data['PC2'],
c=[colors[cluster_id]],
label=f'Cluster {cluster_id}',
s=180, alpha=0.7, edgecolors='black', linewidth=0.5)
# 添加任务简称
for idx, row in cluster_data.iterrows():
task_id = str(idx)
plt.annotate(task_id,
(row['PC1'], row['PC2']),
xytext=(6, 6),
textcoords='offset points',
fontsize=8,
alpha=0.9,
bbox=dict(boxstyle='round,pad=0.2', facecolor='white', alpha=0.7))
plt.xlabel(f'PC1 ({pca.explained_variance_ratio_[0]:.1%} variance)', fontsize=12)
plt.ylabel(f'PC2 ({pca.explained_variance_ratio_[1]:.1%} variance)', fontsize=12)
plt.title('MT25 Task Embeddings PCA with K-Means Clustering\n5 Task Groups Identified',
fontsize=14, fontweight='bold')
plt.grid(True, alpha=0.3)
plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left', fontsize=10)
plt.tight_layout()
plt.savefig('./img/mt25_pca_clustered.png', dpi=300, bbox_inches='tight')
plt.close()
print("\n✓ 聚类PCA图已保存到: ./img/mt25_pca_clustered.png")
# 12. 额外:保存相似度热力图
plt.figure(figsize=(14, 12))
# 使用简称创建热力图
short_names = [t.replace('-v3', '') for t in task_names]
sns.heatmap(similarity_matrix,
xticklabels=short_names,
yticklabels=short_names,
cmap='viridis',
center=0.5,
square=True,
cbar_kws={'shrink': 0.8})
plt.title('MT25 Task Similarity Matrix\n(Cosine Similarity)',
fontsize=14, fontweight='bold')
plt.xticks(rotation=45, ha='right', fontsize=9)
plt.yticks(rotation=0, fontsize=9)
plt.tight_layout()
plt.savefig('./img/mt25_similarity_heatmap.png', dpi=300, bbox_inches='tight')
plt.close()
print("✓ 相似度热力图已保存到: ./img/mt25_similarity_heatmap.png")
print("\n" + "="*70)
print("所有图表已成功保存到 ./img/ 目录")
print("="*70)