def get_data_distribution(dataset=[]): labels = [] for data in dataset: for _,label in data: labels.extend(tf.argmax(label,axis=1).numpy()) y, idx, count = tf.unique_with_counts(labels) return y, count def visualize_data(class_data,count): plt.figure(figsize=(10,10)) plt.bar(class_data,count,align="center") plt.tight_layout() plt.title("Class Distribution") plt.ylabel("Number of Data") plt.xlabel("Class in number") plt.show() __ __