Skip to content

Instantly share code, notes, and snippets.

@leo-pfeiffer
Last active May 12, 2021 16:24
Show Gist options
  • Select an option

  • Save leo-pfeiffer/d25adb26395d6085ecdc307a8f0b279b to your computer and use it in GitHub Desktop.

Select an option

Save leo-pfeiffer/d25adb26395d6085ecdc307a8f0b279b to your computer and use it in GitHub Desktop.
Calculate the gini index and entropy for a list of classes
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