- refacored state
This commit is contained in:
@@ -0,0 +1,86 @@
|
||||
import numpy as np
|
||||
from numpy.random import uniform
|
||||
|
||||
|
||||
def create_test_state() -> np.array:
|
||||
return np.array([['1', '2', '3'], ['4', '5', '6'], ['7', '8', '9']])
|
||||
|
||||
|
||||
def create_empty_state() -> np.array:
|
||||
return np.array([['-', '-', '-'], ['-', '-', '-'], ['-', '-', '-']])
|
||||
|
||||
|
||||
def softmax(values: np.array, eps=1.0E-6) -> np.array:
|
||||
result = (values + eps) / (sum(values + eps))
|
||||
return result
|
||||
|
||||
|
||||
def sample(values: np.array, with_debug=False):
|
||||
# Normalize and sort
|
||||
probs = softmax(values)
|
||||
sorted_indices = np.argsort(probs)
|
||||
sorted_probs = probs[sorted_indices]
|
||||
z = uniform()
|
||||
if with_debug:
|
||||
print(f"Probs={probs}")
|
||||
print(f"Z={z:.3f}")
|
||||
index = None
|
||||
p_sum = 0
|
||||
for idx, p in enumerate(sorted_probs):
|
||||
p_sum += p
|
||||
if z <= p_sum:
|
||||
index = sorted_indices[idx]
|
||||
break
|
||||
|
||||
return index
|
||||
|
||||
|
||||
def test_sample(data):
|
||||
p = np.array(data)
|
||||
c = np.array([0]*len(data))
|
||||
|
||||
for i in range(0, 1000):
|
||||
index = sample(p, with_debug=False)
|
||||
c[index] += 1
|
||||
|
||||
print(c)
|
||||
|
||||
|
||||
def to_state_string(state, state_nex=None):
|
||||
sp = ' '
|
||||
sp_arrow = ' => '
|
||||
sp_ = [sp, sp_arrow, sp]
|
||||
|
||||
def col(str_in, state):
|
||||
str_out = str_in
|
||||
for c in range(0, 3):
|
||||
char = state[r][c]
|
||||
if char == '-':
|
||||
char = ' '
|
||||
str_out += "|" + char
|
||||
str_out += '|'
|
||||
return str_out
|
||||
|
||||
state_str = ''
|
||||
for r in range(0, 3):
|
||||
if state is not None:
|
||||
state_str = col(state_str, state)
|
||||
|
||||
if state_nex is not None:
|
||||
state_str += sp_[r]
|
||||
state_str = col(state_str, state_nex)
|
||||
|
||||
if r != 2:
|
||||
state_str += '\x0A'
|
||||
|
||||
return state_str
|
||||
|
||||
print("Testing")
|
||||
test_sample([0.7, 0.1, 0.1, 0.1])
|
||||
test_sample([0.1, 0.1, 0.2, 0.1])
|
||||
test_sample([0.2, 0.8])
|
||||
test_sample([0.1, 0.1])
|
||||
test_sample([0.3, 0.0, 0.2, 0.0])
|
||||
test_sample([0.0, 0.0, 0.0, 0.0])
|
||||
test_sample([0.5, 0.0, 0.2, 0.3])
|
||||
test_sample([0.3, 0.1, 0.4, 0.2])
|
||||
Reference in New Issue
Block a user