Find Reachable State Space from a Transition Model Function
Source:R/find_reachable_states.R
find_reachable_states.RdFinds the reachable state space from a transition model function
that takes the arguments model, action, start.state and
returns a named vector with the probabilities for the resulting end.states
(typically only the ones with a probability greater than 1).
Usage
find_reachable_states(
transition_model,
start_state,
actions,
model = NULL,
horizon = Inf,
progress = TRUE
)Arguments
- start_state
labels of the start states.
- actions
a vector with the available actions.
- model
if needed, the model passed on to the transition model function.
- horizon
only return states reachable in the given horizon.
- progress
logical; show a progress bar?
- transition_model
a transition function (see details for requirements).
Details
The function performs a (depth-limited) depth-first traversal of the search state space and returns a vector with the names of all encountered states. This vector can be used as the states for creating an MDP model.
The transition function needs to be a function with the argument list
model, action, start.state which returns named vector only containing
the non-zero probabilities named by the corresponding end state. A partial model
that contains actions and start.state will be supplied to the function.
Examples
# define an MDP for Tic-Tac-Toe
# state description: matrix with the characters _, x, and o
# can be converted into a label of 9 characters
# set of actions
A <- as.character(1:9)
# helper functions
ttt_empty_board <- function() matrix('_', ncol = 3, nrow = 3)
ttt_state2label <- function(state) paste(state, collapse = '')
ttt_label2state <- function(label) matrix(strsplit(label, "")[[1]],
nrow = 3, ncol = 3)
ttt_available_actions <- function(state) {
if (length(state) == 1L) state <- ttt_label2state(state)
which(state == "_")
}
ttt_result <- function(state, player, action) {
if (length(state) == 1L) state <- ttt_label2state(state)
if (state[action] != "_")
stop("Illegal action.")
state[action] <- player
state
}
ttt_terminal <- function(state) {
if (length(state) == 1L) state <- ttt_label2state(state)
# Check the board for a win and return one of
# 'x', 'o', 'd' (draw), or 'n' (for next move)
win_possibilities <- rbind(state,
t(state),
diag(state),
diag(t(state)))
wins <- apply(win_possibilities, MARGIN = 1, FUN = function(x) {
if (x[1] != '_' && length(unique(x)) == 1) x[1]
else '_'
})
if (any(wins == 'x'))
return('x')
if (any(wins == 'o'))
return('o')
# Check for draw
if (sum(state == '_') < 1)
return('d')
return('n')
}
# define the transition function:
# * return a probability vector for an action in a start state
# * we define the special states 'win', 'loss', and 'draw'
P <- function(model, action, start.state) {
action <- as.integer(action)
# absorbing states
if (start.state %in% c('win', 'loss', 'draw', 'illegal')) {
return(structure(1, names = start.state))
}
# avoid illegal action by going to the very expensive illegal state
if (!(action %in% ttt_available_actions(start.state))) {
return(structure(1, names = "illegal"))
}
# make x's move
next_state <- ttt_result(start.state, 'x', action)
# terminal?
term <- ttt_terminal(next_state)
if (term == 'x') {
return(structure(1, names = "win"))
} else if (term == 'o') {
return(structure(1, names = "loss"))
} else if (term == 'd') {
return(structure(1, names = "draw"))
}
# it is o's turn
actions_of_o <- ttt_available_actions(next_state)
possible_end_states <- lapply(
actions_of_o,
FUN = function(a)
ttt_result(next_state, 'o', a)
)
# fix terminal states
term <- sapply(possible_end_states, ttt_terminal)
possible_end_states <- sapply(possible_end_states, ttt_state2label)
possible_end_states[term == 'x'] <- 'win'
possible_end_states[term == 'o'] <- 'loss'
possible_end_states[term == 'd'] <- 'draw'
possible_end_states <- unique(possible_end_states)
return(structure(rep(1 / length(possible_end_states),
length(possible_end_states)),
names = possible_end_states))
}
# define the reward
R <- rbind(
R_( value = 0),
R_(end.state = 'win', value = +1),
R_(end.state = 'loss', value = -1),
R_(end.state = 'draw', value = +.5),
R_(end.state = 'illegal', value = -Inf),
# Note: there is no more reward once the agent is in a terminal state
R_(start.state = 'win', value = 0),
R_(start.state = 'loss', value = 0),
R_(start.state = 'draw', value = 0),
R_(start.state = 'illegal', value = 0)
)
# start state
start <- ttt_state2label(ttt_empty_board())
start
#> [1] "_________"
# find the reachable state space
S <- union(c('win', 'loss', 'draw', 'illegal'),
reachable_states(P, start_state = start, actions = A))
head(S)
#> [1] "win" "loss" "draw" "illegal" "xoxo_ox__" "xoxoo_xxo"
tictactoe <- MDP(S, A, P, R, discount = 1, start = start, name = "TicTacToe")
tictactoe
#> MDPModel, MDP - TicTacToe
#> Discount factor: 1
#> Horizon: Inf epochs
#> Size: 9 actions / 2527 states
#> Storage: transition prob as function / reward as data.frame. Total size: 274.6 Kb
#> Start: _________
#> Model list components: ‘name’, ‘discount’, ‘horizon’, ‘states’,
#> ‘actions’, ‘start’, ‘transition_model’, ‘reward’, ‘info’