Add files via upload

This commit is contained in:
费政聪
2024-07-18 14:20:29 +08:00
committed by GitHub
parent d2a932ab0b
commit 85136ff419
4 changed files with 307 additions and 0 deletions
+69
View File
@@ -0,0 +1,69 @@
import os
import json
import torch
from models import selected_ids_list
from download import find_model
from models import DiT_models
from diffusion import create_diffusion
def image_class_expert_ratio():
image_size = 256
model = "DiT-S/2"
num_classes = 1000
device = "cuda"
ckpt_path = "results/002-DiT-S-2/checkpoints/ckpt.pt"
num_sampling_steps = 250
cfg_scale = 1.5
every_class_sample = 50
torch.manual_seed(1234)
torch.set_grad_enabled(False)
latent_size = image_size // 8
model = DiT_models[model](
input_size=latent_size,
num_classes=num_classes,
num_experts=8,
num_experts_per_tok=2,
).to(device)
if ckpt_path is not None:
print('load from: ', ckpt_path)
state_dict = find_model(ckpt_path)
model.load_state_dict(state_dict)
model.eval()
diffusion = create_diffusion(str(num_sampling_steps))
for i in range(1000):
experts_ids = []
for j in range(every_class_sample):
class_labels = [i]
# Create sampling noise:
n = len(class_labels)
z = torch.randn(n, 4, latent_size, latent_size, device=device)
y = torch.tensor(class_labels, device=device)
# Setup classifier-free guidance:
z = torch.cat([z, z], 0)
y_null = torch.tensor([1000] * n, device=device)
y = torch.cat([y, y_null], 0)
model_kwargs = dict(y=y, cfg_scale=cfg_scale)
# Sample images:
samples = diffusion.p_sample_loop(
model.forward_with_cfg, z.shape, z, clip_denoised=False, model_kwargs=model_kwargs, progress=True, device=device
)
print(i, j)
print(len(selected_ids_list), len(selected_ids_list[0]), len(selected_ids_list[0][0]))
tmp_ids_list = selected_ids_list[-3000:]
print(len(tmp_ids_list), len(tmp_ids_list[0]), len(tmp_ids_list[0][0]))
print(tmp_ids_list[0][0])
experts_ids.append(tmp_ids_list)
#break
#continue
print(len(experts_ids))
tgt_path = os.path.join('experts', str(i)+'.json')
with open(tgt_path, 'w') as f:
json.dump(experts_ids, f,)
+86
View File
@@ -0,0 +1,86 @@
import json
def normalization(data):
_range = np.max(data) - np.min(data)
return (data - np.min(data)) / _range
def heatmap_plt(static_list):
import numpy as np
import matplotlib.pyplot as plt
y = [x * 100 for x in range(10)]
y1 = [x * 10 for x in range(10)]
for i in range(12):
data = [x[i] for x in static_list]
data = np.stack(data, axis=0)
# data = normalization(data)
plt.subplot(2, 6, i + 1)
plt.imshow(data, cmap='Oranges', aspect='auto')
if i == 0 or i==6:
plt.ylabel('Image classes')
plt.title('MoE Layer ' + str(i))
if i < 6:
plt.xticks([])
if i== 0 or i == 6:
plt.yticks(y1, y)
else:
plt.yticks([])
# plt.colorbar() # 添加颜色条
# plt.title('Value Heatmap Example')
# plt.xlabel('Expert ids')
# plt.ylabel('Image classes')
plt.subplots_adjust(wspace=0.1)
plt.show()
import numpy as np
def static_for_class(data, layer_num=12):
static_list = np.zeros((layer_num, 8)) # [[0] * 8] * layer_num
for data_list in data:
for j in range(3000):
row = j % layer_num
for k in range(256):
col = k % 8
expert_id = data_list[j][k][0]
expert_id2 = data_list[j][k][1]
static_list[row][expert_id] += 1
static_list[row][expert_id2] += 1
return static_list
import os
calculate_flag = False
static_list = []
for i in range(0, 1000):
if i % 10 != 0:
continue
path = os.path.join('data/experts', str(i) + '.json')
tgt_path = os.path.join('data/class', str(i) + '.npy')
if calculate_flag == True:
with open(path, 'r') as f:
data_list = json.load(f)
# print(len(data_list), len(data_list[0]), len(data_list[0][0]))
# 50, 3000 (250 * 12), 512 (256 * 2), 2
# print(data_list[0][0][0])
# print(data_list[0][1][0])
# print(data_list[0][2][0])
# continue
static = static_for_class(data_list)
np.save(tgt_path, static)
else:
static = np.load(tgt_path)
print(i)
print(static)
static_list.append(static)
# print(static_list)
heatmap_plt(static_list)
+79
View File
@@ -0,0 +1,79 @@
import json
def normalization(data):
_range = np.max(data) - np.min(data)
return (data - np.min(data)) / _range
def heatmap_plt(static_list):
import numpy as np
import matplotlib.pyplot as plt
#y = [x * 100 for x in range(10)]
#y1 = [x * 10 for x in range(10)]
for i in range(12):
# data = [x[i] for x in static_list]
data = static_list[i] #np.stack(data, axis=0)
data = normalization(data)
plt.subplot(2, 6, i + 1)
plt.imshow(data, cmap='Greens', aspect='auto')
if i == 0 or i==6:
plt.ylabel('Image patch')
plt.title('MoE Layer ' + str(i))
if i < 6:
plt.xticks([])
if i== 0 or i == 6:
continue
else:
plt.yticks([])
# plt.colorbar() # 添加颜色条
# plt.title('Value Heatmap Example')
# plt.xlabel('Expert ids')
# plt.ylabel('Image classes')
plt.subplots_adjust(wspace=0.1)
plt.show()
import numpy as np
import os
calculate_flag = False
static_list = np.zeros((12, 256, 8)) # (expert layer, step, expert id)
tgt_path = os.path.join('data/patch', 'patch.npy')
if calculate_flag == True:
for i in range(0, 1000):
if i % 10 != 0:
continue
print(i)
path = os.path.join('data/experts', str(i) + '.json')
with open(path, 'r') as f:
data_list1 = json.load(f)
# print(len(data_list), len(data_list[0]), len(data_list[0][0]))
# 50, 3000 (250 * 12), 512 (256 * 2), 2
# print(data_list1[0][0][0])
# print(data_list1[0][1][0])
# print(data_list1[0][2][0])
# static = static_for_class(data_list)
for data_list in data_list1:
for j in range(3000):
row = j % 12
for k in range(256):
expert_id = data_list[j][k][0]
expert_id2 = data_list[j][k][1]
static_list[row][k][expert_id] += 1
static_list[row][k][expert_id2] += 1
np.save(tgt_path, static_list)
else:
static_list = np.load(tgt_path)
print(static_list.shape)
heatmap_plt(static_list)
+73
View File
@@ -0,0 +1,73 @@
import json
def normalization(data):
_range = np.max(data) - np.min(data)
return (data - np.min(data)) / _range
def heatmap_plt(static_list):
import numpy as np
import matplotlib.pyplot as plt
for i in range(12):
# data = [x[i] for x in static_list]
data = static_list[i] #np.stack(data, axis=0)
data = normalization(data)
plt.subplot(2, 6, i + 1)
plt.imshow(data, cmap='Blues', aspect='auto')
if i == 0 or i==6:
plt.ylabel('Inference step')
plt.title('MoE Layer ' + str(i))
if i < 6:
plt.xticks([])
if i== 0 or i == 6:
continue
else:
plt.yticks([])
# plt.colorbar() # 添加颜色条
# plt.title('Value Heatmap Example')
# plt.xlabel('Expert ids')
# plt.ylabel('Image classes')
plt.subplots_adjust(wspace=0.1)
plt.show()
import numpy as np
import os
calculate_flag = False
static_list = np.zeros((12, 250, 8)) # (expert layer, step, expert id)
tgt_path = os.path.join('data/step', 'step.npy')
if calculate_flag == True:
for i in range(0, 1000):
print(i)
path = os.path.join('data/experts', str(i) + '.json')
with open(path, 'r') as f:
data_list1 = json.load(f)
# print(len(data_list), len(data_list[0]), len(data_list[0][0]))
# 50, 3000 (250 * 12), 512 (256 * 2), 2
# print(data_list1[0][0][0])
# print(data_list1[0][1][0])
# print(data_list1[0][2][0])
# static = static_for_class(data_list)
for data_list in data_list1:
for j in range(3000):
row = j % 12
row2 = j // 12
for k in range(256):
col = k % 8
expert_id = data_list[j][k][0]
expert_id2 = data_list[j][k][1]
static_list[row][row2][expert_id] += 1
static_list[row][row2][expert_id2] += 1
np.save(tgt_path, static_list)
else:
static_list = np.load(tgt_path)
print(static_list.shape)
heatmap_plt(static_list)