Files
Mujoco_WASM/python/mujoco/experimental/studio/web/state_payload.h
T
Matija Kecman 84950fa371 MuJoCo Web Viewer: Implement parallel chunked model downloading and in-place model reloading.
This change introduces parallel chunked downloading for large model files (.mjb) directly into WASM linear memory, enables model loading without full page reloads and reduces memory overhead.

Key changes:
- Implement a chunked model endpoint in the Python web server to support range requests.
- Add parallel chunked fetching in the frontend with retry logic and a single-fetch fallback.
- Increase initial WASM memory to 3 GB to accommodate large models and prevent heap fragmentation.
- Support in-place model reloading in the C++ client, including texture cache invalidation when the Filament context is recreated.
- Display a model download progress bar and model parsing/loading banner to the UI.
- Fixes model drag and drop (caused by typo in sessionId, corrected to session_id).

PiperOrigin-RevId: 960249236
Change-Id: Icac89e6a4ca099882aaf9b111744c1b6ab0cc6c1
2026-08-06 05:51:43 -07:00

144 lines
5.8 KiB
C++

// Copyright 2026 DeepMind Technologies Limited
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// https://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
// This file defines the serialization format for the web viewer's browser
// client render payload containing the data needed so that the browser can
// render the scene using the following call:
//
// Render(model, data, perturb, camera, vis_options, width, height, extra_geoms)
//
// The arguments come from the Python process:
//
// * model : fetched once over HTTP as /model; its runtime-mutable
// parts (opt/vis/stat) are re-sent in the render state block.
// * data : streamed as the physics state vector (mjSTATE_INTEGRATION);
// the browser recomputes the rest via mj_setState/mj_forward.
// * width/height: the browser canvas size.
// * extra_geoms : optional variable-size kTagExtraGeoms block.
// * ... : the rest of the arguments are sent as a fixed-size block
//
// The payload (SerializeStatePayload) is a sequence of tagged blocks:
//
// [StatePayloadHeader][u32 tag][u32 size][payload]...
//
// The payload is serialized by Python, sent over the /state WebSocket, and
// parsed by the browser.
//
// TODO(matijak): Try shrinking the physics block: float32 (or quantized) values
// instead of doubles, and/or delta-encoding against the client's last-acked
// payload. The /state ack (web_server.py) tells the server which snapshot each
// client last applied, which is the baseline that delta compression needs. For
// 100humanoids.xml the payload is ~181 KB of doubles and dominates slow links.
#ifndef MUJOCO_PYTHON_EXPERIMENTAL_STUDIO_WEB_STATE_PAYLOAD_H_
#define MUJOCO_PYTHON_EXPERIMENTAL_STUDIO_WEB_STATE_PAYLOAD_H_
#include <cstddef>
#include <cstdint>
#include <vector>
#include <mujoco/mujoco.h>
namespace mujoco::studio {
// "MJWS" as little-endian bytes. This magic constant identifies the
// StateServer WebSocket payload header and helps detect malformed or
// misrouted messages.
constexpr uint32_t kStatePayloadMagic =
'M' | ('J' << 8) | ('W' << 16) | ('S' << 24);
constexpr uint16_t kStatePayloadVersion = 1;
struct StatePayloadHeader {
uint32_t magic = kStatePayloadMagic;
uint16_t version = kStatePayloadVersion;
uint16_t nblocks = 0;
// CRC32 of the model's MJB bytes. When this changes, the browser must
// refetch /model before applying any further state.
uint32_t model_crc32 = 0;
};
static_assert(sizeof(StatePayloadHeader) == 12);
// Block tags. Readers must skip unknown tags.
enum StateBlockTag : uint32_t {
kTagPhysicsState = 1, // [i32 mjtState spec signature][mjtNum values...]
kTagRenderState = 2, // fixed-size block of kRenderStateSize bytes
kTagExtraGeoms = 3, // n x mjvGeom (n = size / sizeof(mjvGeom))
};
struct StateBlockHeader {
uint32_t tag = 0;
uint32_t size = 0;
};
static_assert(sizeof(StateBlockHeader) == 8);
// Fixed byte size of the render state block appended after physics state.
// These are plain C structs of int/float/double members whose total size is
// fixed, independent of the model and generally negligible compared to the size
// of the physics state
constexpr size_t kRenderStateSize =
sizeof(mjvCamera) + sizeof(mjvPerturb) + sizeof(mjvOption) +
sizeof(mjOption) + sizeof(mjVisual) + sizeof(mjStatistic) + mjNRNDFLAG;
// Maximum number of extra geoms serialized per frame. Bounds the shared
// memory buffer the StateServer allocates; WebViewer truncates longer lists.
constexpr uint32_t kMaxExtraGeoms = 1024;
// Upper bound of a serialized payload, used to size the StateServer's shared
// memory buffer. `physics_bytes` is mj_stateSize(...) * sizeof(mjtNum).
size_t MaxStatePayloadSize(size_t physics_bytes);
// Serialize the complete state payload sent over the state WebSocket.
std::vector<std::byte> SerializeStatePayload(
uint32_t model_crc32, int32_t physics_spec, const void* physics,
size_t physics_bytes, const mjvCamera& camera, const mjvPerturb& perturb,
const mjvOption& vis_options, const mjOption& opt, const mjVisual& vis,
const mjStatistic& stat, const std::vector<uint8_t>& render_flags,
const mjvGeom* extra_geoms, size_t extra_geom_count);
// Parsed view into a serialized payload. Pointers alias the input buffer and
// are NOT guaranteed to be aligned; so you must memcpy the data out before use.
struct StatePayloadView {
uint32_t model_crc32 = 0;
int32_t physics_spec = 0;
const std::byte* physics = nullptr;
size_t physics_bytes = 0;
const std::byte* render_state = nullptr; // kRenderStateSize bytes when non-null
const std::byte* extra_geoms = nullptr; // extra_geom_count * sizeof(mjvGeom)
size_t extra_geom_count = 0;
};
// Parses a payload produced by SerializeStatePayload. Returns false if the
// buffer is malformed (bad magic/version or out-of-bounds block). Blocks
// with unknown tags are skipped.
bool ParseStatePayload(const void* data, size_t size, StatePayloadView* out);
// A decoded render state block.
struct RenderStateView {
mjvCamera camera;
mjvPerturb perturb;
mjvOption vis_options;
mjOption opt;
mjVisual vis;
mjStatistic stat;
uint8_t render_flags[mjNRNDFLAG];
};
// Decodes a render state block produced by the serializer. `data` must hold
// kRenderStateSize bytes (e.g. StatePayloadView::render_state).
void ParseRenderState(const std::byte* data, RenderStateView* out);
} // namespace mujoco::studio
#endif // MUJOCO_PYTHON_EXPERIMENTAL_STUDIO_WEB_STATE_PAYLOAD_H_