Skip to content

Instantly share code, notes, and snippets.

@corajr
Last active December 20, 2015 11:49
Show Gist options
  • Select an option

  • Save corajr/6126303 to your computer and use it in GitHub Desktop.

Select an option

Save corajr/6126303 to your computer and use it in GitHub Desktop.
Genre classification test
#!/usr/bin/python
# -*- coding: utf-8 -*-
import csv
import os
from collections import defaultdict, Counter
comparisons = defaultdict(dict)
true_genres = {}
test_genres = {}
genre_confusion = defaultdict(Counter)
def nn(trackname):
neighbors = comparisons[trackname].keys()
return min(neighbors, key=lambda x: comparisons[trackname][x])
def parse_csv(filename, parse_func):
with file(filename) as f:
dialect = csv.Sniffer().sniff(f.read(1024))
f.seek(0)
reader = csv.reader(f, dialect)
for row in reader:
parse_func(row)
# parse CSVs
def parse_comparison_row(row):
comparisons[row[0]][row[1]] = float(row[4])
comparisons[row[1]][row[0]] = float(row[4])
def parse_tracklist_row(row):
true_genres[os.path.basename(row[5])] = row[0]
def parse_groundtruth_row(row):
true_genres[os.path.basename(row[0])] = row[1]
parse_csv('comparisons.csv', parse_comparison_row)
parse_csv('tracklist.csv', parse_tracklist_row)
parse_csv('ground_truth.csv', parse_groundtruth_row)
for track in true_genres.keys():
try:
track_nn = nn(track)
except SystemExit, KeyboardInterrupt:
sys.exit(1)
except:
print track
true_genre = true_genres[track]
test_genre = true_genres[track_nn]
test_genres[track] = test_genre
genre_confusion[true_genre][test_genre] += 1
# assess accuracy
correct = 0
all_items = 0
for (genre, genre_row) in genre_confusion.iteritems():
correct += genre_confusion[genre][genre]
all_items += sum(genre_row.values())
print '{:02%} correct'.format(float(correct) / all_items)
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment