mirror of
https://github.com/rohitg00/ai-engineering-from-scratch.git
synced 2026-10-02 01:54:39 +08:00
Revert "chore(phase-09): scrub in-prose banned reference-repo mentions"
This reverts commit 236198f6da.
This commit is contained in:
@@ -48,12 +48,12 @@ Use the same 4×4 GridWorld from Lesson 01. We add a stochastic variant: with pr
|
||||
SLIP = 0.1
|
||||
|
||||
def transitions(state, action):
|
||||
if state == TERMINAL:
|
||||
return [(state, 0.0, 1.0)]
|
||||
outcomes = []
|
||||
for direction, prob in action_probs(action):
|
||||
outcomes.append((apply_move(state, direction), -1.0, prob))
|
||||
return outcomes
|
||||
if state == TERMINAL:
|
||||
return [(state, 0.0, 1.0)]
|
||||
outcomes = []
|
||||
for direction, prob in action_probs(action):
|
||||
outcomes.append((apply_move(state, direction), -1.0, prob))
|
||||
return outcomes
|
||||
```
|
||||
|
||||
`transitions(s, a)` returns a list of `(s', r, p)`. This is the entire model.
|
||||
@@ -64,17 +64,17 @@ Given a policy `π(s) = {action: prob}`, iterate the Bellman equation until `V`
|
||||
|
||||
```python
|
||||
def policy_evaluation(policy, gamma=0.99, tol=1e-6):
|
||||
V = {s: 0.0 for s in states()}
|
||||
while True:
|
||||
delta = 0.0
|
||||
for s in states():
|
||||
v = sum(pi_a * sum(p * (r + gamma * V[s_prime])
|
||||
for s_prime, r, p in transitions(s, a))
|
||||
for a, pi_a in policy(s).items())
|
||||
delta = max(delta, abs(v - V[s]))
|
||||
V[s] = v
|
||||
if delta < tol:
|
||||
return V
|
||||
V = {s: 0.0 for s in states()}
|
||||
while True:
|
||||
delta = 0.0
|
||||
for s in states():
|
||||
v = sum(pi_a * sum(p * (r + gamma * V[s_prime])
|
||||
for s_prime, r, p in transitions(s, a))
|
||||
for a, pi_a in policy(s).items())
|
||||
delta = max(delta, abs(v - V[s]))
|
||||
V[s] = v
|
||||
if delta < tol:
|
||||
return V
|
||||
```
|
||||
|
||||
### Step 3: policy improvement
|
||||
@@ -83,28 +83,28 @@ Replace `π` with the greedy policy w.r.t. `V`. If `π` did not change, return
|
||||
|
||||
```python
|
||||
def policy_improvement(V, gamma=0.99):
|
||||
new_policy = {}
|
||||
for s in states():
|
||||
best_a = max(
|
||||
ACTIONS,
|
||||
key=lambda a: sum(p * (r + gamma * V[s_prime])
|
||||
for s_prime, r, p in transitions(s, a)),
|
||||
)
|
||||
new_policy[s] = best_a
|
||||
return new_policy
|
||||
new_policy = {}
|
||||
for s in states():
|
||||
best_a = max(
|
||||
ACTIONS,
|
||||
key=lambda a: sum(p * (r + gamma * V[s_prime])
|
||||
for s_prime, r, p in transitions(s, a)),
|
||||
)
|
||||
new_policy[s] = best_a
|
||||
return new_policy
|
||||
```
|
||||
|
||||
### Step 4: stitch them together
|
||||
|
||||
```python
|
||||
def policy_iteration(gamma=0.99):
|
||||
policy = {s: "up" for s in states()} # arbitrary start
|
||||
for _ in range(100):
|
||||
V = policy_evaluation(lambda s: {policy[s]: 1.0}, gamma)
|
||||
new_policy = policy_improvement(V, gamma)
|
||||
if new_policy == policy:
|
||||
return V, policy
|
||||
policy = new_policy
|
||||
policy = {s: "up" for s in states()} # arbitrary start
|
||||
for _ in range(100):
|
||||
V = policy_evaluation(lambda s: {policy[s]: 1.0}, gamma)
|
||||
new_policy = policy_improvement(V, gamma)
|
||||
if new_policy == policy:
|
||||
return V, policy
|
||||
policy = new_policy
|
||||
```
|
||||
|
||||
Typical convergence on 4×4: 4–6 outer iterations. Outputs `V*(0,0) ≈ -6` and a policy that strictly decreases the step count.
|
||||
@@ -113,19 +113,19 @@ Typical convergence on 4×4: 4–6 outer iterations. Outputs `V*(0,0) ≈ -6` an
|
||||
|
||||
```python
|
||||
def value_iteration(gamma=0.99, tol=1e-6):
|
||||
V = {s: 0.0 for s in states()}
|
||||
while True:
|
||||
delta = 0.0
|
||||
for s in states():
|
||||
v = max(sum(p * (r + gamma * V[s_prime])
|
||||
for s_prime, r, p in transitions(s, a))
|
||||
for a in ACTIONS)
|
||||
delta = max(delta, abs(v - V[s]))
|
||||
V[s] = v
|
||||
if delta < tol:
|
||||
break
|
||||
policy = policy_improvement(V, gamma)
|
||||
return V, policy
|
||||
V = {s: 0.0 for s in states()}
|
||||
while True:
|
||||
delta = 0.0
|
||||
for s in states():
|
||||
v = max(sum(p * (r + gamma * V[s_prime])
|
||||
for s_prime, r, p in transitions(s, a))
|
||||
for a in ACTIONS)
|
||||
delta = max(delta, abs(v - V[s]))
|
||||
V[s] = v
|
||||
if delta < tol:
|
||||
break
|
||||
policy = policy_improvement(V, gamma)
|
||||
return V, policy
|
||||
```
|
||||
|
||||
Same fixed point, fewer lines of code.
|
||||
|
||||
Reference in New Issue
Block a user