Last active
May 12, 2021 16:24
-
-
Save leo-pfeiffer/d25adb26395d6085ecdc307a8f0b279b to your computer and use it in GitHub Desktop.
Calculate the gini index and entropy for a list of classes
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| import argparse | |
| from decimal import * | |
| from functools import reduce, partial | |
| import math | |
| def _class_count(dic, x): | |
| """ | |
| Callback function to count elements in a list. | |
| """ | |
| if x in dic: | |
| dic[x] += Decimal(1) | |
| else: | |
| dic[x] = Decimal(1) | |
| return dic | |
| def _gini_sum(total, x, y): | |
| """ | |
| Callback function to calculate one term of the sum in the gini index. | |
| """ | |
| return x + (y / total) ** 2 | |
| def _entropy_sum(total, x, y): | |
| """ | |
| Callback function to calculate one term of the sum in the entropy. | |
| """ | |
| p_i = (y / total) | |
| return x + p_i * Decimal(math.log(p_i, 2)) | |
| def calculate_gini(classes): | |
| """ | |
| Calculate the gini index for a list of classes. | |
| """ | |
| class_counts = reduce(_class_count, classes, {}) | |
| total_sum = sum(class_counts.values()) | |
| gini_sum = partial(_gini_sum, total_sum) | |
| gini = 1 - reduce(gini_sum, class_counts.values(), 0) | |
| return class_counts, total_sum, gini | |
| def calculate_entropy(classes): | |
| """ | |
| Calculate the entropy for a list of classes. | |
| """ | |
| class_counts = reduce(_class_count, classes, {}) | |
| total_sum = sum(class_counts.values()) | |
| entropy_sum = partial(_entropy_sum, total_sum) | |
| gini = (-1) * reduce(entropy_sum, class_counts.values(), 0) | |
| return class_counts, total_sum, gini | |
| def calculate_gini_gain(classes, true_values): | |
| """ | |
| Calculate the gini gain for a list of classes with corresponding labels. | |
| """ | |
| assert len(classes) == len(true_values) | |
| # number of classes | |
| n = Decimal(len(classes)) | |
| # initialise_values | |
| gini_base = calculate_gini(true_values)[2] | |
| gini_overall = 0 | |
| result = {'gini': {}, 'gini_base': gini_base, 'gini_overall': None} | |
| for c in list(set(classes)): | |
| class_n = Decimal(sum([x==c for x in classes])) | |
| # indices corresponding to the class | |
| c_index = [i for i in range(int(n)) if classes[i] == c] | |
| # class proportion | |
| class_p = class_n / n | |
| # labels resulting from the split | |
| split_labels = [true_values[i] for i in c_index] | |
| # calculate gini | |
| gini = calculate_gini(split_labels)[2] | |
| result['gini'][c] = {'prop': class_p, 'gini': gini, 'count': class_n} | |
| gini_overall += class_p * gini | |
| result['gini_overall'] = gini_overall | |
| return result | |
| def calculate_entropy_gain(classes, true_values): | |
| """ | |
| Calculate the entropy gain for a list of classes with corresponding labels. | |
| """ | |
| assert len(classes) == len(true_values) | |
| # number of classes | |
| n = Decimal(len(classes)) | |
| # initialise_values | |
| entropy_base = calculate_entropy(true_values)[2] | |
| entropy_overall = 0 | |
| result = {'entropy': {}, 'entropy_base': entropy_base, 'entropy_overall': None} | |
| for c in list(set(classes)): | |
| class_n = Decimal(sum([x==c for x in classes])) | |
| # indices corresponding to the class | |
| c_index = [i for i in range(int(n)) if classes[i] == c] | |
| # class proportion | |
| class_p = class_n / n | |
| # labels resulting from the split | |
| split_labels = [true_values[i] for i in c_index] | |
| # calculate entropy | |
| entropy = calculate_entropy(split_labels)[2] | |
| result['entropy'][c] = {'prop': class_p, 'entropy': entropy, 'count': class_n} | |
| entropy_overall += class_p * entropy | |
| result['entropy_overall'] = entropy_overall | |
| return result | |
| if __name__ == '__main__': | |
| getcontext().prec = 4 | |
| # example usage: | |
| # simple gini and entropy calculation for one attribute: | |
| # python purity.py -c A A A A B B B B B B B B C C C C C C C C -g -e | |
| # info gain and gini gain calculation for an attribute and corresponding labels: | |
| # python purity.py -c t t f f t t f f t -t 1 1 0 1 0 0 0 1 0 -ee -gg | |
| # -g : Calculate gini index | |
| # -e : Calculate entropy | |
| # -c : List of attribute values (interpreted as string) | |
| # -t : "True values", i.e. the labels (must be integer) | |
| # -gg: Calculate gini gain (requires specification of -t) | |
| # -ee: Calculate entropy (requires specification of -t) | |
| cli=argparse.ArgumentParser() | |
| cli.add_argument("-c", "--classes", nargs="*", type=str, default=[]) | |
| cli.add_argument("-t", "--true_values", nargs="*", type=int, default=[]) | |
| cli.add_argument("-g", "--gini", nargs="?", type=bool, const=True) | |
| cli.add_argument("-e", "--entropy", nargs="?", type=bool, const=True) | |
| cli.add_argument("-gg", "--gini_gain", nargs="?", type=bool, const=True) | |
| cli.add_argument("-ee", "--entropy_gain", nargs="?", type=bool, const=True) | |
| args = cli.parse_args() | |
| classes = args.classes | |
| true_values = args.true_values | |
| if args.gini: | |
| print("===== Gini =====\n") | |
| gini = calculate_gini(classes) | |
| print(f"Class counts: {gini[0]}\nSum: {gini[1]}\nGini: {gini[2]}\n") | |
| if args.entropy: | |
| print("===== Entropy =====\n") | |
| entr = calculate_entropy(classes) | |
| print(f"Class counts: {entr[0]}\nSum: {entr[1]}\nEntropy: {entr[2]}\n") | |
| if args.gini_gain: | |
| print("===== Gini gain =====\n") | |
| gg = calculate_gini_gain(classes, true_values) | |
| print(f"Gini Base: {gg['gini_base']}") | |
| print(f"Gini Overall: {gg['gini_overall']}") | |
| print(f"Gini Gain: {gg['gini_base'] - gg['gini_overall']}\n") | |
| for c in list(set(classes)): | |
| print(f"{c} ===\nProp: {gg['gini'][c]['prop']}") | |
| print(f"Count: {gg['gini'][c]['count']}") | |
| print(f"Gini: {gg['gini'][c]['gini']}\n") | |
| if args.entropy_gain: | |
| print("===== Entropy gain =====\n") | |
| gg = calculate_entropy_gain(classes, true_values) | |
| print(f"Entropy Base: {gg['entropy_base']}") | |
| print(f"Entropy Overall: {gg['entropy_overall']}") | |
| print(f"Info Gain: {gg['entropy_base'] - gg['entropy_overall']}\n") | |
| for c in list(set(classes)): | |
| print(f"{c} ===\nProp: {gg['entropy'][c]['prop']}") | |
| print(f"Count: {gg['entropy'][c]['count']}") | |
| print(f"Entropy: {gg['entropy'][c]['entropy']}\n") | |
| exit(0) |
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment