Kimhi's picture
Upload StreamPETR EVA-02 source without weights
6a176cb verified
Raw
History Blame Contribute Delete
3.53 kB
import torch
class TemporalMemoryController:
def __init__(
self,
enabled=False,
keyframe_interval=1,
updater='reuse',
updater_cfg=None,
train_with_teacher=False,
distill_weight=1.0,
inference_teacher_forbidden=True,
profiler=None,
debug=False,
):
if keyframe_interval < 1:
raise ValueError('keyframe_interval must be >= 1')
self.enabled = enabled
self.keyframe_interval = keyframe_interval
self.updater_name = updater.get('type', updater.get('name', 'reuse')) if isinstance(updater, dict) else updater
self.train_with_teacher = train_with_teacher
self.inference_teacher_forbidden = inference_teacher_forbidden
self.debug = debug
self.profiler = profiler
self.frame_index = 0
self.scene_token = None
self.memory = None
self.last_distillation_loss = None
from .updaters import build_memory_updater
self.updater = build_memory_updater(updater, **(updater_cfg or {}))
from .distillation import DistillationLoss
self.distillation_loss = DistillationLoss(weight=distill_weight)
def reset(self, reason='manual'):
if self.debug:
print(f'[temporal_memory] reset reason={reason}')
self.frame_index = 0
self.memory = None
def get_or_update(
self,
extract_full,
scene_token=None,
prev_exists=None,
partial_current=None,
ego_motion_embedding=None,
extract_teacher=None,
training=False,
):
self.last_distillation_loss = None
if not self.enabled:
return extract_full()
if scene_token != self.scene_token:
self.scene_token = scene_token
self.reset(reason='scene_boundary')
if prev_exists is not None and not bool(prev_exists):
self.reset(reason='discontinuity')
is_keyframe = (
self.memory is None
or self.keyframe_interval == 1
or self.frame_index % self.keyframe_interval == 0
)
if is_keyframe:
if self.debug:
print(f'[temporal_memory] frame={self.frame_index} decision=keyframe')
self.memory = extract_full()
output = self.memory
else:
if self.debug:
print(f'[temporal_memory] frame={self.frame_index} decision=non_keyframe updater={self.updater_name}')
if self.profiler is None:
output = self.updater(
self.memory,
partial_current=partial_current,
ego_motion_embedding=ego_motion_embedding)
else:
with self.profiler.measure('updater'):
output = self.updater(
self.memory,
partial_current=partial_current,
ego_motion_embedding=ego_motion_embedding)
if not training and extract_teacher is not None and self.inference_teacher_forbidden:
raise AssertionError('Full-feature teacher is forbidden on non-keyframe inference.')
if training and self.train_with_teacher and extract_teacher is not None:
with torch.no_grad():
teacher = extract_teacher()
self.last_distillation_loss = self.distillation_loss(output, teacher)
self.frame_index += 1
return output