diff --git a/analysis/expert_data.py b/analysis/expert_data.py new file mode 100644 index 0000000..3df8584 --- /dev/null +++ b/analysis/expert_data.py @@ -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,) + \ No newline at end of file diff --git a/analysis/heatmap_class.py b/analysis/heatmap_class.py new file mode 100644 index 0000000..e66e4b2 --- /dev/null +++ b/analysis/heatmap_class.py @@ -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) + diff --git a/analysis/heatmap_patch.py b/analysis/heatmap_patch.py new file mode 100644 index 0000000..fe396c2 --- /dev/null +++ b/analysis/heatmap_patch.py @@ -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) + diff --git a/analysis/heatmap_step.py b/analysis/heatmap_step.py new file mode 100644 index 0000000..f0893eb --- /dev/null +++ b/analysis/heatmap_step.py @@ -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) +