LINE Solver
MATLAB API documentation
Loading...
Searching...
No Matches
rl_td_agent.m
1classdef rl_td_agent < handle
2 properties
3 v; % value function
4 Q; % Q function
5 vSize; % size of value function
6 QSize; % size of Q function
7 epsilon = 1; % explore-exploit rate
8 eps_decay = 0.99; % explore-exploit rate decay
9 lr = 0.05; % learning rate
10 end
11
12 methods
13 function obj=rl_td_agent(lr, eps, epsDecay)
14 obj.lr = lr;
15 obj.epsilon = eps;
16 obj.eps_decay = epsDecay;
17 obj.v = 0;
18 obj.vSize = 0;
19 obj.Q = 0;
20 obj.QSize = 0;
21 end
22
23 function reset(obj, env)
24 obj.v = 0;
25 obj.vSize = 0;
26 obj.Q = 0;
27 obj.QSize = 0;
28 env.reset();
29 end
30
31 function v = getValueFunction(obj)
32 v = obj.v;
33 end
34
35 function Q = getQFunction(obj)
36 Q = obj.Q;
37 end
38
39 % see _kb/03-api-layer.md for rationale
40
41 function solve(obj, env)
42 obj.reset(env);
43
44 obj.v = zeros((zeros(1, env.actionSize)+env.stateSize + 5)); % value function
45 obj.Q = rand([(zeros(1, env.actionSize)+env.stateSize + 5), env.actionSize]); % Q function
46 obj.vSize = size(obj.v);
47 obj.QSize = size(obj.Q);
48
49 x = zeros(1, env.actionSize); % initial state
50 n = zeros(1, env.actionSize); % initial previous state
51 % t_prev = 0; % time of last event
52 t = 0; % time of current event
53 % dt = 0; % time period between two successive events
54 c = 0; % incurred costs between the visits
55 T = 0; % total discounted elapsed time
56 C = 0; % total discounted costs
57
58 num_episodes = 1e4;
59 eps = obj.epsilon;
60 j = 0;
61
62 while j < num_episodes
63 if mod(j, 1e3)==0
64 line_printf('[%s] running episode #%d .\n',mfilename,j);
65 end
66
67 % if mod(j, 100) == 0
68 eps = eps * obj.eps_decay;
69 % end
70
71 % t_prev = t;
72 [dt, depNode] = env.sample(); % how to successive sampling
73 t = dt + t;
74 % t = dt + t_prev;
75
76 c = c + sum(x) * dt;
77
78 if ismember(depNode, env.idxOfSourceInNodes) % new job
79 if env.isInActionSpace(env.model.nodes)
80 % create an exploit-explore policy
81 next_locs = zeros(env.actionSize, env.actionSize) + x + 1 + eye(env.actionSize);
82 next_states = obj.get_state_from_locs(obj.vSize, next_locs);
83 policy = obj.createGreedyPolicy(obj.v(next_states), eps, env.actionSize);
84
85 action = sum(rand >= cumsum([0, policy]));
86 else
87 action = find(x==min(x)); % JSQ
88 if length(action)>1
89 action = randomsample(action, 1);
90 end
91 end
92
93 x(action) = x(action) + 1;
94 env.update(x);
95 % see _kb/03-api-layer.md for rationale
96
97 elseif ismember(depNode, env.idxOfQueueInNodes) % dep from Queue, idx: Node{depNode}
98 x(env.idxOfQueueInNodes == depNode) = max(0, x(env.idxOfQueueInNodes == depNode) - 1);
99 env.update(x);
100 % State.afterEvent(sn, ind, inspace, event, class, isSimulation)
101 end
102
103 if env.isInStateSpace(env.model.nodes)
104 j = j + 1;
105 T = env.gamma * T + t;
106 C = env.gamma * C + c;
107 mean_cost_rate = C/T;
108
109 prev_state = obj.get_state_from_loc(obj.vSize, n+1);
110 cur_state = obj.get_state_from_loc(obj.vSize, x+1);
111 obj.v(prev_state) = (1-obj.lr)*obj.v(prev_state) + obj.lr*(c - t*mean_cost_rate + obj.v(cur_state)); % here "obj.v(cur_state) * env.gamma" ?
112 obj.v = obj.v - obj.v(1);
113
114 t = 0;
115 c = 0;
116 n = x;
117 end
118
119 end
120 end
121
122 function s = get_state_from_locs(obj, objSize, locs)
123 s = zeros(1, size(locs,1));
124 for i=1:size(locs,1)
125 s(i) = obj.get_state_from_loc(objSize, locs(i,:));
126 end
127 end
128 end
129
130 methods(Static)
131 function policy = createGreedyPolicy(state_Q, epsilon, nA)
132 policy = ones(1, nA) * epsilon / nA;
133 argmin = find(state_Q-min(state_Q)<GlobalConstants.FineTol);
134 policy(argmin) = policy(argmin) + (1-epsilon)/length(argmin);
135 end
136
137 function s = get_state_from_loc(objSize, loc)
138 s = 0;
139 if size(objSize,2) == size(loc, 2)
140 for i=1:size(objSize,2)
141 if i==1
142 s = s + loc(i);
143 else
144 s = s + (loc(i)-1) * prod(objSize(1:(i-1)));
145 end
146 end
147 end
148 end
149
150 end
151
152end
153
154