FloodDiffusion2-Live / window_control.py
caiyiyi1998's picture
Initial commit
9a25493
Raw History Blame Contribute Delete
6.33 kB
"""Replan a pending path from its last committed spring state."""
import copy
import math
from pathlib import Path
import sys
sys.path.insert(0, str(Path(__file__).resolve().parent / 'space'))
from control_loop import SpringPath, effective_slot
class _DampedSpringPath(SpringPath):
"""Smooth target velocity before the existing position-following spring."""
def __init__(self, frequency=10., input_response_seconds=.35):
super().__init__(frequency)
self.input_response_seconds = float(input_response_seconds)
self.input_velocity = [0., 0.]
def step(self, *, x, z, shift, slot, speeds):
selected = effective_slot(slot, shift)
dt = 1/30.
previous = self.position.copy()
old_heading = self.heading
if self.frame:
norm = max(1., math.hypot(x,z))
desired = [x/norm*speeds[selected], z/norm*speeds[selected]]
velocity_decay = math.exp(-dt/self.input_response_seconds)
for axis in (0,1):
offset = self.input_velocity[axis]-desired[axis]
self.target[axis] += desired[axis]*dt+offset*self.input_response_seconds*(1-velocity_decay)
self.input_velocity[axis] = desired[axis]+offset*velocity_decay
decay = math.exp(-self.frequency*dt)
for axis in (0,1):
offset = self.position[axis]-self.target[axis]
c = self.velocity[axis]+self.frequency*offset
self.position[axis] = self.target[axis]+(offset+c*dt)*decay
self.velocity[axis] = (self.velocity[axis]-self.frequency*c*dt)*decay
dx,dz = self.position[0]-previous[0], self.position[1]-previous[1]
if math.hypot(dx,dz)>1e-5:
target_heading = math.atan2(dx,dz)
difference = (target_heading-self.heading+math.pi)%(2*math.pi)-math.pi
self.heading += max(-math.pi*dt,min(math.pi*dt,difference*(1-math.exp(-10*dt))))
c,s = math.cos(self.heading),math.sin(self.heading)
row = [self.heading-old_heading,c*dx-s*dz,s*dx+c*dz]
else:
row = [0.,0.,0.]
point = dict(frame=self.frame,x=self.position[0],z=self.position[1],
heading=self.heading,slot=selected,base_slot=int(slot),run=selected==1,
target_x=self.target[0],target_z=self.target[1])
self.frame += 1
return row,point
def _clone_spring(spring):
"""Copy the complete SpringPath state without recursive object traversal."""
clone = type(spring).__new__(type(spring))
clone.frequency = spring.frequency
clone.frame = spring.frame
clone.heading = spring.heading
clone.target = spring.target.copy()
clone.position = spring.position.copy()
clone.velocity = spring.velocity.copy()
if isinstance(spring, _DampedSpringPath):
clone.input_response_seconds = spring.input_response_seconds
clone.input_velocity = spring.input_velocity.copy()
return clone
class WindowPath:
def __init__(self):
self.committed = _DampedSpringPath()
self.states = []
self.points = []
self.rows = []
def plan(self, text_metas, *, x, z, shift, slot, speeds, version):
"""Hold the latest controller input throughout the uncommitted horizon."""
spring = _clone_spring(self.committed)
states, points, rows = [], [], []
selected = effective_slot(slot, shift)
for original in text_metas:
row, point = spring.step(x=x, z=z, shift=shift, slot=slot, speeds=speeds)
if point['frame'] != original['frame']:
raise RuntimeError('Path/text frame alignment was lost')
# The caller supplies the text condition actually used for each frame.
point.update(slot=original['slot'], base_slot=original['base_slot'],
run=original['run'], prompt=original['prompt'],
text_version=original['text_version'], version=version,
path_slot=selected)
states.append(_clone_spring(spring))
points.append(point)
rows.append(row)
self.states, self.points, self.rows = states, points, rows
return rows, points
def commit_first(self):
if not self.states:
raise RuntimeError('No pending plan to commit')
self.committed = self.states.pop(0)
return self.rows.pop(0), self.points.pop(0)
def target_point(self, slot):
spring = self.states[-1] if self.states else self.committed
return spring.target_point(slot)
def check_planner():
import math
speeds = [0.8, 2.5, 0.45, 0.8]
original, planner = _DampedSpringPath(), WindowPath()
# Under constant input, replanning must reproduce exactly the existing path.
pending = []
reference = []
max_error = 0.
for step in range(380):
row, point = original.step(x=0.6, z=0.8, shift=False, slot=0, speeds=speeds)
reference.append(row)
pending.append(dict(frame=step, slot=0, base_slot=0, run=False,
prompt='walk', text_version=0))
rows, points = planner.plan(pending, x=0.6, z=0.8, shift=False,
slot=0, speeds=speeds, version=0)
start = planner.committed.frame
max_error = max(max_error, max(abs(a-b) for i,r in enumerate(rows)
for a,b in zip(r, reference[start+i])))
if step >= 29:
planner.commit_first()
pending.pop(0)
assert max_error == 0., max_error
frozen = copy.deepcopy(planner.committed.__dict__)
rows, _ = planner.plan(pending, x=-1., z=0., shift=False, slot=0,
speeds=speeds, version=1)
assert planner.committed.__dict__ == frozen
assert any(abs(a-b) > 1e-8 for a,b in zip(rows[0],reference[planner.committed.frame]))
planner.commit_first()
assert planner.committed.frame == frozen['frame']+1
assert all(math.isfinite(v) for row in rows for v in row)
return {'constant_input_max_error': max_error, 'committed_history_unchanged': True,
'next_frame_path_revised': True, 'frames_checked': 380}
if __name__ == '__main__':
print(check_planner())