forked from camall3n/visgrid
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathutils.py
More file actions
69 lines (60 loc) · 2.16 KB
/
Copy pathutils.py
File metadata and controls
69 lines (60 loc) · 2.16 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
import argparse
import glob
import json
import numpy as np
import os
import random
from sklearn.neighbors import KernelDensity
#import torch
def manhattan_dist(pos1, pos2):
x1, y1 = pos1
x2, y2 = pos2
return np.abs(x2 - x1) + np.abs(y2 - y1)
def get_parser():
"""Return a nicely formatted argument parser
This function is a simple wrapper for the argument parser I like to use,
which has a stupidly long argument that I always forget.
"""
return argparse.ArgumentParser(formatter_class=argparse.ArgumentDefaultsHelpFormatter)
def fit_kde(x, bw=0.03):
p = KernelDensity(bandwidth=bw, kernel='tophat')
p.fit(x)
return p
def MI(x, y):
xy = np.concatenate([x, y], axis=-1)
log_pxy = fit_kde(xy).score_samples(xy)
log_px = fit_kde(x).score_samples(x)
log_py = fit_kde(y).score_samples(y)
log_ratio = log_pxy - log_px - log_py
return np.mean(log_ratio)
def load_experiment(tag, coefs=None):
logfiles = sorted(glob.glob(os.path.join('results/logs', tag + '*', 'train-*.txt')))
seeds = [f.split('-')[-1].split('.')[0] for f in logfiles]
logs = [open(f, 'r').read().splitlines() for f in logfiles]
def read_log(log, coefs=coefs):
results = [json.loads(item) for item in log]
fields = results[0].keys()
data = dict([(f, np.asarray([item[f] for item in results])) for f in fields])
if coefs is None:
coefs = {
'L_inv': 1.0,
'L_fwd': 0.1,
'L_cpc': 1.0,
'L_fac': 0.1,
}
if 'L' not in fields:
data['L'] = sum([
coefs[f] * data[f] if f != 'L_fac' else coefs[f] * (data[f] - 1)
for f in coefs.keys()
])
return data
results = [read_log(log) for log in logs]
data = dict(zip(seeds, results))
return data
def get_good_color(color):
colorname = color
colorname = 'gold' if colorname == 'yellow' else colorname
colorname = 'c' if colorname == 'cyan' else colorname
colorname = 'm' if colorname == 'magenta' else colorname
colorname = 'silver' if colorname in ['gray', 'grey'] else colorname
return colorname