diff --git a/positronic/server/positronic_server.py b/positronic/server/positronic_server.py index b65e84e34..377429b89 100644 --- a/positronic/server/positronic_server.py +++ b/positronic/server/positronic_server.py @@ -100,6 +100,17 @@ async def cache_rerun_assets(request: Request, call_next): return response +def _make_serializable(obj): + """Ensure obj is JSON serializable (e.g. convert datetime).""" + if isinstance(obj, datetime): + return obj.isoformat() + if isinstance(obj, dict): + return {k: _make_serializable(v) for k, v in obj.items()} + if isinstance(obj, list): + return [_make_serializable(v) for v in obj] + return obj + + def _iter_file_chunks(path: str, *, chunk_size: int = 128 * 1024): with open(path, 'rb') as source: while True: @@ -185,16 +196,6 @@ async def episode_viewer(request: Request, episode_id: int): size_mb = meta.get('size_mb') size_mb_display = f'{size_mb:.2f}' if isinstance(size_mb, int | float) else None - # Ensure static_data is JSON serializable (e.g. handle datetime) - def _make_serializable(obj): - if isinstance(obj, datetime): - return obj.isoformat() - if isinstance(obj, dict): - return {k: _make_serializable(v) for k, v in obj.items()} - if isinstance(obj, list): - return [_make_serializable(v) for v in obj] - return obj - return templates.TemplateResponse( 'episode.html', { @@ -401,6 +402,28 @@ async def api_dataset_status(): } +@app.get('/api/episode/{episode_id}') +@require_dataset +async def api_episode(episode_id: int): + ds = app_state.get('dataset') + try: + episode = ds[episode_id] + except IndexError as e: + raise HTTPException(status_code=404, detail='Episode not found') from e + + meta = episode.meta + size_mb = meta.get('size_mb') + + return { + 'episode_id': episode_id, + 'num_episodes': len(ds), + 'task': episode.static.get('task', None), + 'episode_path': meta.get('path'), + 'episode_size_mb': f'{size_mb:.2f}' if isinstance(size_mb, int | float) else None, + 'static_data': _make_serializable(episode.static), + } + + @app.get('/api/episode_rrd/{episode_id}') @require_dataset async def api_episode_rrd(episode_id: int): diff --git a/positronic/server/static/app.js b/positronic/server/static/app.js index 05e28b0c8..3d62d9747 100644 --- a/positronic/server/static/app.js +++ b/positronic/server/static/app.js @@ -522,6 +522,7 @@ document.addEventListener('DOMContentLoaded', () => { function initializeSidebar(staticData) { const sidebarContent = document.querySelector('.sidebar-content-wrapper tbody'); + sidebarContent.innerHTML = ''; function isNestable(value) { return typeof value === 'object' && value !== null; diff --git a/positronic/server/static/sw.js b/positronic/server/static/sw.js new file mode 100644 index 000000000..3f204a035 --- /dev/null +++ b/positronic/server/static/sw.js @@ -0,0 +1,31 @@ +// Service Worker to cache rerun WASM and JS assets. +// These files are large (~35MB WASM) and don't change between episodes. + +const CACHE_NAME = 'rerun-assets-v1'; +const RERUN_PATH_PREFIX = '/static/rerun/'; + +self.addEventListener('install', () => self.skipWaiting()); +self.addEventListener('activate', (event) => { + event.waitUntil( + caches.keys().then((names) => + Promise.all(names.filter((n) => n !== CACHE_NAME).map((n) => caches.delete(n))) + ).then(() => self.clients.claim()) + ); +}); + +self.addEventListener('fetch', (event) => { + const url = new URL(event.request.url); + if (!url.pathname.startsWith(RERUN_PATH_PREFIX)) return; + + event.respondWith( + caches.open(CACHE_NAME).then((cache) => + cache.match(event.request).then((cached) => { + if (cached) return cached; + return fetch(event.request).then((response) => { + if (response.ok) cache.put(event.request, response.clone()); + return response; + }); + }) + ) + ); +}); diff --git a/positronic/server/templates/base.html b/positronic/server/templates/base.html index 5995fbb9a..a2f3427b3 100644 --- a/positronic/server/templates/base.html +++ b/positronic/server/templates/base.html @@ -28,6 +28,7 @@ + {% block scripts %}{% endblock %} diff --git a/positronic/server/templates/episode.html b/positronic/server/templates/episode.html index a642f8e8a..8f7f4cb2a 100644 --- a/positronic/server/templates/episode.html +++ b/positronic/server/templates/episode.html @@ -7,48 +7,165 @@ {% endblock %} {% block scripts %} + + - + {% endblock %} @@ -150,31 +260,25 @@ Ep. -
- / {{ num_episodes }}
- {% if task %} : {{ task }}{% endif %} + {% if task %} : {{ task }}{% endif %}
- {% if episode_path or episode_size_mb %} -
- {% if episode_path %} - Episode path: - {{ episode_path }} - {% endif %} - {% if episode_size_mb %} - ({{ episode_size_mb }} MB) - {% endif %} +
+ Episode path: + {{ episode_path or '' }} + {% if episode_size_mb %}({{ episode_size_mb }} MB){% endif %}
- {% endif %}
diff --git a/positronic/vendors/lerobot/train.py b/positronic/vendors/lerobot/train.py index 0a3d234a2..d4756b259 100644 --- a/positronic/vendors/lerobot/train.py +++ b/positronic/vendors/lerobot/train.py @@ -21,7 +21,7 @@ from lerobot.configs.policies import PreTrainedConfig from lerobot.configs.train import TrainPipelineConfig from lerobot.envs.configs import EnvConfig, FeatureType, PolicyFeature -from lerobot.policies.xvla.configuration_xvla import XVLAConfig # noqa: F401 — registers policy choices +from lerobot.policies.smolvla.configuration_smolvla import SmolVLAConfig # noqa: F401 — registers policy from lerobot.utils.constants import ACTION, OBS_IMAGES, OBS_STATE from positronic import utils @@ -93,15 +93,14 @@ def _update_config(cfg: TrainPipelineConfig, **cfg_kwargs): raise AttributeError(f'Could not update config for {k}') from e -@cfn.config(codec=lerobot_codecs.ee, base_model='lerobot/smolvla_base', num_train_steps=None) -@pos3.with_mirror() -def train( +def _train( input_path: str, exp_name: str, output_dir: str, codec: Codec, base_model: str, num_train_steps: int | None, + batch_size: int, **cfg_kwargs, ): if isinstance(codec, str): @@ -134,6 +133,7 @@ def train( job_name=exp_name, eval_freq=0, log_freq=10, + batch_size=batch_size, steps=num_train_steps if num_train_steps is not None else 100_000, ) @@ -144,6 +144,9 @@ def train( _update_config(cfg, **cfg_kwargs) + if 'policy.scheduler_decay_steps' not in cfg_kwargs: + cfg.policy.scheduler_decay_steps = cfg.steps - cfg.policy.scheduler_warmup_steps + if cfg.resume: checkpoints_dir = Path(cfg.output_dir) / 'checkpoints' if not checkpoints_dir.exists(): @@ -182,9 +185,23 @@ def train( logging.info('Training finished.') +@cfn.config(codec=lerobot_codecs.ee, base_model='lerobot/smolvla_base', num_train_steps=None, batch_size=64) +@pos3.with_mirror() +def train(input_path, exp_name, output_dir, codec, base_model, num_train_steps, batch_size, **cfg_kwargs): + _train(input_path, exp_name, output_dir, codec, base_model, num_train_steps, batch_size, **cfg_kwargs) + + +@cfn.config(codec=lerobot_codecs.ee, base_model='lerobot/smolvla_base', num_train_steps=None, batch_size=64) +@pos3.with_mirror() +def full_finetune(input_path, exp_name, output_dir, codec, base_model, num_train_steps, batch_size, **cfg_kwargs): + cfg_kwargs.setdefault('policy.freeze_vision_encoder', False) + cfg_kwargs.setdefault('policy.train_expert_only', False) + _train(input_path, exp_name, output_dir, codec, base_model, num_train_steps, batch_size, **cfg_kwargs) + + def _internal_main(): init_logging() - cfn.cli(train) + cfn.cli({'train': train, 'full_finetune': full_finetune}) if __name__ == '__main__': diff --git a/pyproject.toml b/pyproject.toml index fb43baba7..f7cdd7db8 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -22,6 +22,7 @@ dependencies = [ "dearpygui", "dm_control", "fastapi", + "jinja2", # Required by starlette's Jinja2Templates (optional dep of fastapi) "fire", "httpx", "msgpack",