diff --git a/agent_config.yaml b/agent_config.yaml index e6c69dab..e593ada0 100644 --- a/agent_config.yaml +++ b/agent_config.yaml @@ -1,14 +1,14 @@ # Agent Configuration agent: - # Select the model provider. + # Select the model provider. # Local: 'ollama' # Remote: 'openrouter', 'google', 'openai' (requires API tokens in .env) provider: ollama - + # Local Settings ollama_model: qwen2.5:3b-instruct fallback_to_local: true - + # Remote Model Selection openrouter_model: mistralai/mistral-small-3.1-24b-instruct:free google_model: gemini-1.5-flash diff --git a/requirements.txt b/requirements.txt index 687f0bb0..6b183c4d 100644 --- a/requirements.txt +++ b/requirements.txt @@ -26,7 +26,11 @@ onnxruntime==1.23.2 xxhash==3.6.0 tables==3.10.2 httpx==0.28.1 +<<<<<<< Updated upstream pydantic==2.11.10 +======= +pydantic>=2.6 +>>>>>>> Stashed changes dotenv==0.9.9 h5py==3.15.1 diff --git a/weightslab/__init__.py b/weightslab/__init__.py index 84d57565..42b4d367 100644 --- a/weightslab/__init__.py +++ b/weightslab/__init__.py @@ -41,7 +41,6 @@ __maintainer__ = 'Guillaume PELLUET' __credits__ = 'GrayBox' __license__ = 'BSD 2-clause' - __all__ = [ "watch_or_edit", "serve", diff --git a/weightslab/backend/cli.py b/weightslab/backend/cli.py index 2644a043..6cd577ab 100644 --- a/weightslab/backend/cli.py +++ b/weightslab/backend/cli.py @@ -100,17 +100,29 @@ def _handle_command(cmd: str) -> Any: 'status': 'Show basic status: registered models and optimizers', 'list_models': 'List registered model names in the ledger', 'list_optimizers': 'List registered optimizer names in the ledger', - 'dump': 'Return a sanitized dump of the ledger contents', + 'list_loaders': 'List registered dataloader names in the ledger', + 'list_uids': 'List data sample UIDs. Syntax: list_uids [loader_name] [--discarded] [--limit N]', + 'dump': 'Return a sanitized dump of the ledger contents', 'operate': 'Edit model architecture. Syntax: operate [] ', - 'plot_model': 'Show ASCII tree of model architecture. Syntax: plot_model []', + 'plot_model': 'Show ASCII tree of model architecture. Syntax: plot_model []', + 'discard': 'Discard data samples. Syntax: discard [uid2 ...] [--loader loader_name]', + 'undiscard': 'Un-discard data samples. Syntax: undiscard [uid2 ...] [--loader loader_name]', + 'add_tag': 'Add tag to data sample. Syntax: add_tag [--loader loader_name]', 'hp / hyperparams': 'List or show hyperparameters. Syntax: hp -> list, hp -> show', - # editing hyperparameters is disabled in the CLI (read-only) 'quit / exit': 'Close the client connection' }, 'hyperparams_examples': { 'list': 'hp', 'show': 'hp fashion_mnist', 'set': "set_hp # e.g. set_hp fashion_mnist data.train_loader.batch_size 32", + }, + 'data_examples': { + 'list all UIDs': 'list_uids', + 'list specific loader': 'list_uids train_loader', + 'list discarded only': 'list_uids --discarded', + 'discard samples': 'discard sample_001 sample_002', + 'undiscard samples': 'undiscard sample_001', + 'add tag': 'add_tag sample_001 difficult', } } @@ -180,9 +192,120 @@ def _handle_command(cmd: str) -> Any: if verb == 'list_dataloaders': return {'ok': True, 'dataloaders': GLOBAL_LEDGER.list_dataloaders()} + if verb in ('list_loaders', 'loaders'): + return {'ok': True, 'loaders': GLOBAL_LEDGER.list_dataloaders()} + if verb == 'list_optimizers': return {'ok': True, 'optimizers': GLOBAL_LEDGER.list_optimizers()} + if verb in ('list_uids', 'uids', 'samples'): + # Syntax: list_uids [loader_name] [--discarded] [--limit N] + loader_name = None + show_discarded_only = False + limit = None + + # Parse arguments + i = 1 + while i < len(parts): + if parts[i] == '--discarded': + show_discarded_only = True + elif parts[i] == '--limit' and i + 1 < len(parts): + try: + limit = int(parts[i + 1]) + i += 1 + except ValueError: + return {'ok': False, 'error': 'Invalid limit value'} + elif not parts[i].startswith('--'): + loader_name = parts[i] + i += 1 + + try: + loaders_to_check = [loader_name] if loader_name else GLOBAL_LEDGER.list_dataloaders() + + result = {'ok': True, 'uids': {}} + + for lname in loaders_to_check: + try: + loader = GLOBAL_LEDGER.get_dataloader(lname) + # Unwrap proxy + if hasattr(loader, 'get') and callable(loader.get): + loader = loader.get() + + if loader is None: + continue + + # Try to get dataset from loader + dataset = None + if hasattr(loader, 'dataset'): + dataset = loader.dataset + + if dataset is None: + continue + + # Get UIDs and discard status + uids_list = [] + + # Try different methods to get UIDs + if hasattr(dataset, 'get_sample_uids'): + # Method to get all UIDs + all_uids = dataset.get_sample_uids() + elif hasattr(dataset, 'sample_ids'): + all_uids = dataset.sample_ids + elif hasattr(dataset, 'uids'): + all_uids = dataset.uids + elif hasattr(dataset, '__len__'): + # Fallback: generate UIDs from indices + all_uids = [f"sample_{i:06d}" for i in range(len(dataset))] + else: + all_uids = [] + + # Check discard status for each UID + for uid in all_uids: + is_discarded = False + + # Try to get discard status + if hasattr(dataset, 'is_discarded'): + try: + is_discarded = dataset.is_discarded(uid) + except Exception: + pass + elif hasattr(dataset, 'discarded_samples'): + is_discarded = uid in dataset.discarded_samples + + # Filter based on --discarded flag + if show_discarded_only and not is_discarded: + continue + + # Get tags if available + tags = [] + if hasattr(dataset, 'get_tags'): + try: + tags = dataset.get_tags(uid) + except Exception: + pass + elif hasattr(dataset, 'sample_tags') and hasattr(dataset.sample_tags, 'get'): + tags = dataset.sample_tags.get(uid, []) + + uids_list.append({ + 'uid': uid, + 'discarded': is_discarded, + 'tags': tags + }) + + # Apply limit + if limit and len(uids_list) >= limit: + break + + result['uids'][lname] = uids_list + + except Exception as e: + result['uids'][lname] = {'error': str(e)} + + return result + + except Exception as e: + return {'ok': False, 'error': str(e)} + # Return a lightweight snapshot of all ledger registries if verb in ('ledgers', 'ledger', 'snapshot'): snap = GLOBAL_LEDGER.snapshot() @@ -314,6 +437,156 @@ def _handle_command(cmd: str) -> Any: return {'ok': True, 'operated': True, 'op': (op_type, layer_id, nb), 'model': model_name} + if verb in ('discard', 'undiscard'): + # Syntax: discard [uid2 ...] [--loader loader_name] + # Syntax: undiscard [uid2 ...] [--loader loader_name] + + if len(parts) < 2: + return {'ok': False, 'error': f'usage: {verb} [uid2 ...] [--loader loader_name]'} + + loader_name = None + uids = [] + + # Parse arguments + i = 1 + while i < len(parts): + if parts[i] == '--loader' and i + 1 < len(parts): + loader_name = parts[i + 1] + i += 2 + else: + uids.append(parts[i]) + i += 1 + + if not uids: + return {'ok': False, 'error': 'No UIDs specified'} + + discard_status = 1 if verb == 'discard' else 0 + + try: + # Get loader(s) + if loader_name: + loaders_to_update = {loader_name: GLOBAL_LEDGER.get_dataloader(loader_name)} + else: + # Try all loaders + loader_names = GLOBAL_LEDGER.list_dataloaders() + loaders_to_update = {name: GLOBAL_LEDGER.get_dataloader(name) for name in loader_names} + + results = {'ok': True, 'updated': {}, 'errors': {}} + + for lname, loader in loaders_to_update.items(): + # Unwrap proxy + if hasattr(loader, 'get') and callable(loader.get): + loader = loader.get() + + if loader is None: + continue + + # Get dataset + dataset = getattr(loader, 'dataset', None) if loader else None + + if dataset is None: + continue + + updated_uids = [] + + for uid in uids: + try: + # Try different methods to set discard status + if hasattr(dataset, 'set_discard'): + dataset.set_discard(uid, discard_status) + updated_uids.append(uid) + elif hasattr(dataset, 'discard_sample'): + if discard_status == 1: + dataset.discard_sample(uid) + else: + # undiscard + if hasattr(dataset, 'undiscard_sample'): + dataset.undiscard_sample(uid) + updated_uids.append(uid) + elif hasattr(dataset, 'discarded_samples'): + # Direct set manipulation + if discard_status == 1: + dataset.discarded_samples.add(uid) + else: + dataset.discarded_samples.discard(uid) + updated_uids.append(uid) + else: + results['errors'][uid] = f'No discard method available in {lname}' + except Exception as e: + results['errors'][uid] = str(e) + + if updated_uids: + results['updated'][lname] = updated_uids + + return results + + except Exception as e: + return {'ok': False, 'error': str(e)} + + if verb in ('add_tag', 'tag'): + # Syntax: add_tag [--loader loader_name] + + if len(parts) < 3: + return {'ok': False, 'error': 'usage: add_tag [--loader loader_name]'} + + uid = parts[1] + tag = parts[2] + loader_name = None + + # Parse optional --loader argument + if len(parts) >= 5 and parts[3] == '--loader': + loader_name = parts[4] + + try: + # Get loader(s) + if loader_name: + loaders_to_update = {loader_name: GLOBAL_LEDGER.get_dataloader(loader_name)} + else: + # Try all loaders + loader_names = GLOBAL_LEDGER.list_dataloaders() + loaders_to_update = {name: GLOBAL_LEDGER.get_dataloader(name) for name in loader_names} + + results = {'ok': True, 'updated': {}, 'errors': {}} + + for lname, loader in loaders_to_update.items(): + # Unwrap proxy + if hasattr(loader, 'get') and callable(loader.get): + loader = loader.get() + + if loader is None: + continue + + # Get dataset + dataset = getattr(loader, 'dataset', None) if loader else None + + if dataset is None: + continue + + try: + # Try different methods to add tags + if hasattr(dataset, 'add_tag'): + dataset.add_tag(uid, tag) + results['updated'][lname] = f'Added tag "{tag}" to {uid}' + elif hasattr(dataset, 'sample_tags'): + # Direct dict manipulation + if uid not in dataset.sample_tags: + dataset.sample_tags[uid] = [] + if tag not in dataset.sample_tags[uid]: + dataset.sample_tags[uid].append(tag) + results['updated'][lname] = f'Added tag "{tag}" to {uid}' + else: + results['errors'][lname] = 'No tag method available' + except Exception as e: + results['errors'][lname] = str(e) + + if not results['updated']: + return {'ok': False, 'error': 'Could not add tag to any loader'} + + return results + + except Exception as e: + return {'ok': False, 'error': str(e)} + # Hyperparameters: list / show details and set if verb in ('hp', 'hyperparams'): # hp -> list diff --git a/weightslab/backend/dataloader_interface.py b/weightslab/backend/dataloader_interface.py index b5e3ea3c..63da8f38 100644 --- a/weightslab/backend/dataloader_interface.py +++ b/weightslab/backend/dataloader_interface.py @@ -13,6 +13,7 @@ - global pause control - registration in a global ledger and dynamic batch-size updates based on hyperparameters +- checkpoint-based data loading and reproducible iterator restoration """ import logging from typing import Any, Iterator, Optional @@ -25,9 +26,10 @@ from weightslab.backend.ledgers import ( register_dataloader, get_hyperparams, - resolve_hp_name + resolve_hp_name, + get_checkpoint_manager, ) -from weightslab.utils import filter_kwargs_for_callable +from weightslab.utils import filter_kwargs_for_callable, restore_rng_state # Get Global Logger @@ -64,6 +66,7 @@ def __iter__(self): # Iterate through base sampler, yielding only non-denied indices for idx in self.base_sampler: + # Skip already generated samples based on offset uid = int(self.tracked_dataset.unique_ids[idx]) if uid not in deny_listed_uids: yield idx @@ -80,6 +83,44 @@ def __len__(self): return max(0, total - denied_count) +class OffsetSampler(Sampler): + """A sampler wrapper that skips the first N samples without loading them. + + This enables efficient state restoration by avoiding data preprocessing + for skipped samples. Works with any base sampler. + """ + + def __init__(self, base_sampler: Sampler, offset: int = 0): + """Initialize with a base sampler and number of samples to skip. + + Args: + base_sampler: The underlying sampler to wrap + offset: Number of samples to skip from the beginning + """ + self.base_sampler = base_sampler + self.offset = max(0, offset) + + def __iter__(self): + """Iterate, skipping the first offset samples.""" + iterator = iter(self.base_sampler) + # Skip offset samples without yielding + for _ in range(self.offset): + try: + next(iterator) + except StopIteration: + return + # Yield remaining samples + yield from iterator + + def __len__(self): + """Return remaining length after offset.""" + try: + total = len(self.base_sampler) + return max(0, total - self.offset) + except Exception: + raise TypeError("len not supported for this sampler") + + class MutableBatchSampler: """A simple mutable batch sampler that yields lists of indices. @@ -184,6 +225,7 @@ def __init__( # Internal flags / helpers self._mutable_batch_sampler = None self._dl_build_kwargs: Optional[dict] = None + self._pending_iteration_state: Optional[dict] = None if isinstance(data_loader_or_dataset, DataLoader): logger.warning( @@ -220,6 +262,9 @@ def __init__( (self._reset_iterator, {}) ) + # Load checkpoint data early (before dataloader is used) + self._load_checkpoint_data() + # store kwargs so we can recreate dataloader if needed self._dl_build_kwargs = { "batch_size": batch_size, @@ -299,8 +344,22 @@ def __len__(self): self.init_attributes(self.dataloader) - # Internal iterator used by `_next_batch` - self._iterator: Iterator = iter(self.dataloader) + # Apply pending iteration state if one was loaded from checkpoint + if hasattr(self, '_pending_iteration_state') and self._pending_iteration_state: + try: + self.restore_iteration_state(self._pending_iteration_state) + logger.info(f"Restored dataloader iteration state: {self._pending_iteration_state}") + except Exception as e: + logger.warning(f"Failed to restore pending iteration state: {e}") + finally: + self._pending_iteration_state = None + + # Internal iterator used by `_next_batch` (lazy created to avoid consuming RNG early) + self._iterator: Optional[Iterator] = None + # Track how many samples have been yielded since last reset (for reproducible seeking) + self._samples_yielded: int = 0 + self._sample_offset: int = 0 + self._skipped = [] # Optionally register in the global ledger for cross-thread access. # If no explicit `name` is provided, try to infer a friendly name from @@ -320,6 +379,88 @@ def __len__(self): # Best-effort: ignore registration failures pass + def _load_checkpoint_data(self) -> None: + """Load data checkpoint, RNG state, and dataloader iteration state early. + + This method is called after tracked_dataset initialization to restore + data from the latest checkpoint if available. It: + 1. Loads data snapshot and applies it to the dataframe + 2. Restores RNG state for reproducible shuffling + 3. Restores dataloader iteration state for deterministic seeking + """ + try: + checkpoint_manager = get_checkpoint_manager() + if checkpoint_manager is None: + return + + # Get latest experiment hash + latest_hash = None + if hasattr(checkpoint_manager, 'current_exp_hash') and checkpoint_manager.current_exp_hash: + latest_hash = checkpoint_manager.current_exp_hash + elif hasattr(checkpoint_manager, 'get_latest_hash'): + latest_hash = checkpoint_manager.get_latest_hash() + + if not latest_hash: + return + + # Load checkpoint with data state + checkpoint_data = checkpoint_manager.load_checkpoint( + exp_hash=latest_hash, + load_model=False, + load_weights=False, + load_config=False, + load_data=True, + force=True + ) + + if not checkpoint_data.get('loaded_components'): + return + + # Apply data snapshot to dataframe if loaded + if 'data' in checkpoint_data['loaded_components']: + try: + data_state = checkpoint_data.get('data_state', {}) + snapshot_df = data_state.get('snapshot') + + if snapshot_df is not None and not snapshot_df.empty: + if hasattr(self.tracked_dataset, 'upsert_df'): + self.tracked_dataset.upsert_df(snapshot_df, force_flush=True) + logger.info(f"Applied data snapshot from checkpoint ({len(snapshot_df)} rows)") + except Exception as e: + logger.warning(f"Failed to apply data snapshot: {e}") + + # Restore RNG state for reproducible shuffling + if checkpoint_data.get('rng_state'): + try: + restore_rng_state(checkpoint_data['rng_state']) + logger.debug("Restored RNG state from checkpoint") + except Exception as e: + logger.warning(f"Failed to restore RNG state: {e}") + + # Restore dataloader iteration state for deterministic seeking + if checkpoint_data.get('dataloader_iteration_state'): + try: + iter_state = checkpoint_data['dataloader_iteration_state'] + # Normalize to handle both dict and single state formats + if isinstance(iter_state, dict) and 'samples_yielded' in iter_state: + # Single state format; will be applied when dataloader is ready + self._pending_iteration_state = iter_state + elif isinstance(iter_state, dict): + # Multi-loader format; pick one for this loader + state_for_loader = iter_state.get(self._ledger_name) or iter_state.get('default') or next(iter(iter_state.values()), None) + if state_for_loader: + self._pending_iteration_state = state_for_loader + else: + self._pending_iteration_state = iter_state + + if hasattr(self, '_pending_iteration_state'): + logger.debug(f"Pending iteration state to restore: {self._pending_iteration_state}") + except Exception as e: + logger.warning(f"Failed to parse dataloader iteration state: {e}") + + except Exception as e: + logger.debug(f"Could not load checkpoint data: {e}") + def init_attributes(self, obj): """Expose attributes and methods from the wrapped `obj`. @@ -416,11 +557,17 @@ def __len__(self) -> int: return len(self.dataloader) def __iter__(self) -> Iterator: - """Return an iterator over batches (delegates to the wrapped dataloader).""" + """Return a self-iterating wrapper that auto-resets on exhaustion. + + Returning ``self`` ensures ``__next__`` is used, which already + handles StopIteration by recreating the underlying iterator. This + makes ``for batch in dataloader_interface`` loop forever over epochs + without the user having to call ``reset_iterator`` manually. + """ self._sync_batch_size_from_ledger() - res = iter(self.dataloader) self._wait_if_paused() - return res + self._reset_iterator() # Reset + return self._iterator def __next__(self) -> Any: """Retrieve the next batch; used when iterating directly over the interface.""" @@ -442,7 +589,7 @@ def _sync_batch_size_from_ledger(self) -> None: if hp_name is None: # no hyperparams; optionally use a default try: - self.set_batch_size(64) + self.set_batch_size(1) except RuntimeError: pass return @@ -482,22 +629,167 @@ def _next_batch(self) -> Any: """Return the next batch from the dataloader. If the iterator is exhausted it is automatically reset and iteration - resumes (unless `is_training=False`, in which case StopIteration is - propagated). + resumes. """ try: + if self._iterator is None: + self._reset_iterator() + # # Execute offset + self._execute_offset() + # Generate batch batch = next(self._iterator) + # Count yielded samples to support iteration state capture/restore + self._samples_yielded += 1 except StopIteration: - if not self.is_training: - raise StopIteration("End of dataloader reached.") + # Reset iterator and try again self._reset_iterator() + # Execute offset + self._execute_offset() + # Generate batch batch = next(self._iterator) + self._samples_yielded += 1 return batch + def _execute_offset(self) -> None: + """ + Execute sample offset if set, skipping samples as needed. + This is a fallback mechanism for user-supplied dataloaders where + we cannot use an OffsetSampler. + + TODO (GP): + We can reproduce the random generation of samples by restoring RNG state, if during the previous checkpoints, batchsize changed dynamically and shuffle is True. + """ + if self._sample_offset > 0: + current_bs = self.get_batch_size() + # Fast-forward the iterator by the offset amount + while len(self._skipped) < self._sample_offset: + try: + bs = 4 if self._sample_offset - len(self._skipped) >= 4 else self._sample_offset - len(self._skipped) # Autoscale bs to sample offset + self.set_batch_size(bs) + self._skipped.extend(next(self._iterator)[1].detach().cpu().tolist()) + logger.debug(f"Offset sampler: skipped {len(self._skipped)}/{self._sample_offset} samples: {self._skipped}") + except StopIteration as e: + logger.debug(f"Offset sampler: reached end of iterator while skipping: {e}") + self._reset_iterator() # Reset iterator and try again + + self.set_batch_size(current_bs) + self._skipped = [] + self._sample_offset = 0 + def _reset_iterator(self) -> None: """Reset the internal iterator so `_next_batch()` starts from the beginning.""" self._iterator = iter(self.dataloader) + def reset_iterator(self) -> None: + """Recreate the internal iterator (e.g., after restoring RNG state). + + Call this after restore_rng_state() to get a fresh shuffle with the + restored RNG state: + + rng_state = capture_rng_state() + batch1 = next(dataloader_interface) + restore_rng_state(rng_state) + dataloader_interface.reset_iterator() # Create new iterator with restored RNG + batch1_repeat = next(dataloader_interface) # Same batches! + """ + self._reset_iterator() + + # ------------------------------------------------------------------------- + # Iteration state capture/restore for deterministic resume + # ------------------------------------------------------------------------- + def capture_iteration_state(self) -> dict: + """Capture current iteration position for later restoration. + + Returns a serializable dict that can be stored with checkpoints and + later supplied to `restore_iteration_state` to resume at the same batch + boundary. Works with and without shuffling. When shuffling, ensure + RNG state is also captured/restored before calling `restore_iteration_state`. + """ + return { + "samples_yielded": int(self._samples_yielded), + "batch_size": self.batch_size or 1 + } + + def restore_iteration_state(self, state: dict) -> None: + """Restore iteration position efficiently without reprocessing skipped data. + + For dataloaders we built (with _dl_build_kwargs), this recreates the + dataloader with an OffsetSampler that skips samples at the index level, + avoiding expensive data loading and transforms for skipped batches. + + For shuffled loaders, call this after restoring RNG state. + """ + try: + samples_yielded = int(state.get("samples_yielded", 0)) + batch_size = int(state.get("batch_size", self.batch_size or 1)) + except Exception: + samples_yielded = 0 + batch_size = self.batch_size or 1 + + # Calculate sample offset (how many individual samples to skip) + sample_offset = samples_yielded + + # If we own the dataloader construction, rebuild with offset sampler + if getattr(self, "_dl_build_kwargs", None) is not None and sample_offset > 0: + try: + kwargs = dict(self._dl_build_kwargs) + # Remove kwargs that conflict with using batch_sampler + kwargs.pop("batch_size", None) + shuffle = kwargs.pop("shuffle", False) + num_workers = kwargs.pop("num_workers", 0) + drop_last = kwargs.pop("drop_last", False) + pin_memory = kwargs.pop("pin_memory", False) + collate_fn = kwargs.pop("collate_fn", None) + kwargs.pop("sampler", None) + kwargs.pop("drop_last", None) + kwargs.pop("shuffle", None) + + # Create base sampler + base_sampler = ( + RandomSampler(self.tracked_dataset) + if shuffle + else SequentialSampler(self.tracked_dataset) + ) + + # # Wrap with offset to skip already-yielded samples + # offset_sampler = OffsetSampler(base_sampler, offset=sample_offset) + + # Wrap with masked sampler for deny-listed samples + masked_sampler = MaskedSampler(base_sampler, self.tracked_dataset) + + # Rebuild mutable batch sampler + mbs_cls = type(self._mutable_batch_sampler) if self._mutable_batch_sampler else MutableBatchSampler + mbs = mbs_cls(masked_sampler, batch_size=batch_size, drop_last=drop_last) + self._mutable_batch_sampler = mbs + self._sample_offset = sample_offset + + # Rebuild dataloader with offset sampler + self.dataloader = DataLoader( + self.tracked_dataset, + batch_sampler=mbs, + num_workers=num_workers, + pin_memory=pin_memory, + collate_fn=collate_fn, + # Ensure no conflicting args are passed alongside batch_sampler + **filter_kwargs_for_callable(DataLoader, kwargs) + ) + + # Reset iterator and counter + self._iterator = None + self._samples_yielded = samples_yielded + return + except Exception as e: + logger.warning(f"Failed to restore with offset sampler, falling back to fast-forward: {e}") + + # Fallback: fast-forward approach (less efficient but works for user-supplied dataloaders) + self._reset_iterator() + for _ in range(max(0, samples_yielded)): + try: + next(self._iterator) + self._samples_yielded += 1 + except StopIteration: + break + # ------------------------------------------------------------------------- # Batch-size management # ------------------------------------------------------------------------- diff --git a/weightslab/backend/ledgers.py b/weightslab/backend/ledgers.py index ce1f161c..e7dd6400 100644 --- a/weightslab/backend/ledgers.py +++ b/weightslab/backend/ledgers.py @@ -34,6 +34,8 @@ def __init__(self, obj: Any = None): self._obj = obj def set(self, obj: Any) -> None: + if isinstance(obj, Proxy): + obj = obj.get() self._obj = obj # invalidate any cached iterator when target changes if hasattr(self, '_iterator'): @@ -43,13 +45,19 @@ def set(self, obj: Any) -> None: pass def get(self, default=None) -> Any: - return self._obj if self._obj is not None and default is not None else default + return self._obj if self._obj is not None else default def __getattr__(self, item): - if self._obj is None: + # Use object.__getattribute__ to avoid infinite recursion during unpickling + try: + obj = object.__getattribute__(self, '_obj') + except AttributeError: + raise AttributeError("Proxy target not set") + + if obj is None: raise AttributeError("Proxy target not set") try: - return getattr(self._obj, item) + return getattr(obj, item) except AttributeError: return None @@ -183,16 +191,19 @@ def __init__(self) -> None: self._dataloaders: Dict[str, Any] = {} self._optimizers: Dict[str, Any] = {} self._dataframes: Dict[str, Any] = {} + self._checkpoint_managers: Dict[str, Any] = {} # weak refs self._models_weak: "weakref.WeakValueDictionary[str, Any]" = weakref.WeakValueDictionary() self._dataloaders_weak: "weakref.WeakValueDictionary[str, Any]" = weakref.WeakValueDictionary() self._optimizers_weak: "weakref.WeakValueDictionary[str, Any]" = weakref.WeakValueDictionary() self._dataframes_weak: "weakref.WeakValueDictionary[str, Any]" = weakref.WeakValueDictionary() + self._checkpoint_managers_weak: "weakref.WeakValueDictionary[str, Any]" = weakref.WeakValueDictionary() # proxies mapping name -> Proxy for placeholders self._proxies_models: Dict[str, Proxy] = {} self._proxies_dataloaders: Dict[str, Proxy] = {} self._proxies_optimizers: Dict[str, Proxy] = {} self._proxies_dataframes: Dict[str, Proxy] = {} + self._proxies_checkpoint_managers: Dict[str, Proxy] = {} # hyperparameters registry (name -> dict) self._hyperparams: Dict[str, Dict[str, Any]] = {} self._proxies_hyperparams: Dict[str, Proxy] = {} @@ -458,6 +469,19 @@ def unregister_signal(self, name: str) -> None: self._signals.pop(name, None) self._proxies_signals.pop(name, None) + # Checkpoint managers + def register_checkpoint_manager(self, name: str, manager: Any, weak: bool = False) -> Any: + return self._register(self._checkpoint_managers, self._checkpoint_managers_weak, self._proxies_checkpoint_managers, name, manager, weak=weak) + + def get_checkpoint_manager(self, name: Optional[str] = None) -> Any: + return self._get(self._checkpoint_managers, self._checkpoint_managers_weak, self._proxies_checkpoint_managers, name) + + def list_checkpoint_managers(self) -> List[str]: + return self._list(self._checkpoint_managers, self._checkpoint_managers_weak) + + def unregister_checkpoint_manager(self, name: str) -> None: + self._unregister(self._checkpoint_managers, self._checkpoint_managers_weak, self._proxies_checkpoint_managers, name) + # DataFrames (e.g., shared sample stats managers) def register_dataframe(self, name: str, dataframe: Any, weak: bool = False) -> None: return self._register(self._dataframes, self._dataframes_weak, self._proxies_dataframes, name, dataframe, weak=weak) @@ -519,14 +543,23 @@ def clear(self) -> None: self._dataloaders.clear() self._optimizers.clear() self._dataframes.clear() + self._checkpoint_managers.clear() + self._hyperparams.clear() + self._loggers.clear() + self._signals.clear() self._models_weak.clear() self._dataloaders_weak.clear() self._optimizers_weak.clear() self._dataframes_weak.clear() + self._checkpoint_managers_weak.clear() self._proxies_models.clear() self._proxies_dataloaders.clear() self._proxies_optimizers.clear() self._proxies_dataframes.clear() + self._proxies_checkpoint_managers.clear() + self._proxies_hyperparams.clear() + self._proxies_loggers.clear() + self._proxies_signals.clear() def snapshot(self) -> Dict[str, List[str]]: """Return the current keys for all registries (a lightweight snapshot).""" @@ -538,6 +571,7 @@ def snapshot(self) -> Dict[str, List[str]]: "dataframes": list(self._dataframes.keys()), "hyperparams": list(self._hyperparams.keys()), "loggers": list(self._loggers.keys()), + "checkpoint_managers": list(self._checkpoint_managers.keys()), } def __repr__(self) -> str: @@ -548,7 +582,8 @@ def __repr__(self) -> str: # Module-level singleton GLOBAL_LEDGER = Ledger() -# Convenience top-level wrappers (preserve optional weak param) + +# Model def list_models() -> List[str]: return GLOBAL_LEDGER.list_models() @@ -556,12 +591,16 @@ def register_model(name: str, model: Any, weak: bool = False) -> None: GLOBAL_LEDGER.register_model(name, model, weak=weak) def get_model(name: Optional[str] = None) -> Any: + if name is None: + existing = GLOBAL_LEDGER.list_models() + name = existing[-1] if existing else "main" return GLOBAL_LEDGER.get_model(name) def get_models() -> List[str]: return GLOBAL_LEDGER.list_models() +# Dataloaders def list_dataloaders() -> List[str]: return GLOBAL_LEDGER.list_dataloaders() @@ -569,12 +608,16 @@ def register_dataloader(name: str, dataloader: Any, weak: bool = False) -> None: GLOBAL_LEDGER.register_dataloader(name, dataloader, weak=weak) def get_dataloader(name: Optional[str] = None) -> Any: + if name is None: + existing = GLOBAL_LEDGER.list_dataloaders() + name = existing[-1] if existing else "main" return GLOBAL_LEDGER.get_dataloader(name) def get_dataloaders() -> List[str]: return GLOBAL_LEDGER.list_dataloaders() +# Optimizer def list_optimizers() -> List[str]: return GLOBAL_LEDGER.list_optimizers() @@ -582,16 +625,23 @@ def register_optimizer(name: str, optimizer: Any, weak: bool = False) -> None: GLOBAL_LEDGER.register_optimizer(name, optimizer, weak=weak) def get_optimizer(name: Optional[str] = None) -> Any: + if name is None: + existing = GLOBAL_LEDGER.list_optimizers() + name = existing[-1] if existing else "main" return GLOBAL_LEDGER.get_optimizer(name) def get_optimizers() -> List[str]: return GLOBAL_LEDGER.list_optimizers() +# Hyperparameters def register_hyperparams(name: str, params: Dict[str, Any], weak: bool = False) -> None: GLOBAL_LEDGER.register_hyperparams(name, params, weak=weak) def get_hyperparams(name: Optional[str] = None) -> Any: + if name is None: + existing = GLOBAL_LEDGER.list_hyperparams() + name = existing[-1] if existing else "main" return GLOBAL_LEDGER.get_hyperparams(name) def list_hyperparams() -> List[str]: @@ -625,10 +675,14 @@ def unwatch_hyperparams_file(name: str) -> None: return GLOBAL_LEDGER.unwatch_hyperparams_file(name) +# Logger def register_logger(name: str, logger: Any) -> None: GLOBAL_LEDGER.register_logger(name, logger) def get_logger(name: Optional[str] = None) -> Any: + if name is None: + existing = GLOBAL_LEDGER.list_loggers() + name = existing[-1] if existing else "main" return GLOBAL_LEDGER.get_logger(name) def list_loggers() -> List[str]: @@ -638,10 +692,14 @@ def unregister_logger(name: str) -> None: return GLOBAL_LEDGER.unregister_logger(name) +# Signals def register_signal(name: str, signal: Any) -> None: GLOBAL_LEDGER.register_signal(name, signal) def get_signal(name: Optional[str] = None) -> Any: + if name is None: + existing = GLOBAL_LEDGER.list_signals() + name = existing[-1] if existing else "main" return GLOBAL_LEDGER.get_signal(name) def list_signals() -> List[str]: @@ -651,17 +709,31 @@ def unregister_signal(name: str) -> None: return GLOBAL_LEDGER.unregister_signal(name) +# Checkpoint managers +def register_checkpoint_manager(name: str, manager: Any, weak: bool = False) -> Any: + return GLOBAL_LEDGER.register_checkpoint_manager(name, manager, weak=weak) + +def get_checkpoint_manager(name: Optional[str] = None) -> Any: + if name is None: + existing = GLOBAL_LEDGER.list_checkpoint_managers() + name = existing[-1] if existing else "main" + return GLOBAL_LEDGER.get_checkpoint_manager(name) + +def list_checkpoint_managers() -> List[str]: + return GLOBAL_LEDGER.list_checkpoint_managers() + +def unregister_checkpoint_manager(name: str) -> None: + GLOBAL_LEDGER.unregister_checkpoint_manager(name) + + # DataFrames def register_dataframe(name: str, dataframe: Any, weak: bool = False) -> None: return GLOBAL_LEDGER.register_dataframe(name, dataframe, weak=weak) def get_dataframe(name: Optional[str] = None) -> Any: - """Return a dataframe handle; if none registered, return a Proxy placeholder. - - When name is omitted, default to "sample_stats" so callers get a stable - proxy that will be updated in-place once the dataframe manager registers. - """ - name = name or "sample_stats" + if name is None: + existing = GLOBAL_LEDGER.list_dataframes() + name = existing[-1] if existing else "main" return GLOBAL_LEDGER.get_dataframe(name) def list_dataframes() -> List[str]: @@ -671,6 +743,12 @@ def unregister_dataframe(name: str) -> None: return GLOBAL_LEDGER.unregister_dataframe(name) +def clear_all() -> None: + """Clear all registries (models, dataloaders, optimizers, dataframes, hyperparams, loggers, signals, etc.)""" + return GLOBAL_LEDGER.clear() + + +# Main if __name__ == "__main__": # Quick demonstration import torch diff --git a/weightslab/backend/model_interface.py b/weightslab/backend/model_interface.py index e43fafc0..b12b0535 100644 --- a/weightslab/backend/model_interface.py +++ b/weightslab/backend/model_interface.py @@ -6,7 +6,7 @@ from torch.fx.passes.shape_prop import ShapeProp from torch.fx import symbolic_trace -from weightslab.components.checkpoint import CheckpointManager +from weightslab.components.checkpoint_manager_v2 import CheckpointManagerV2 from weightslab.components.tracking import TrackingMode from weightslab.models.model_with_ops import NetworkWithOps from weightslab.modules.neuron_ops import NeuronWiseOperations @@ -19,6 +19,8 @@ generate_index_maps from weightslab.components.global_monitoring import guard_training_context, guard_testing_context from weightslab.backend.ledgers import get_optimizer, get_optimizers, register_model +from weightslab.backend import ledgers +from weightslab.utils.tools import restore_rng_state # Global logger @@ -76,10 +78,116 @@ def __init__( self.dummy_input = dummy_input.to(device) else: self.dummy_input = th.randn(model.input_shape).to(device) + + # Initialize checkpoint manager and attempt early auto-load before any model-dependent setup + self._checkpoint_manager = None + _checkpoint_auto_every_steps = 0 + _root_log_dir = None + try: + from weightslab.backend.ledgers import list_hyperparams, get_hyperparams + names = list_hyperparams() + chosen = None + if 'main' in names: + chosen = 'main' + elif 'experiment' in names: + chosen = 'experiment' + elif len(names) > 1: + chosen = names[-1] + + if chosen: + hp = get_hyperparams(chosen) + if hasattr(hp, 'get') and not isinstance(hp, dict): + try: + hp = hp.get() + except Exception: + hp = None + if isinstance(hp, dict): + _root_log_dir = hp.get('root_log_dir') or hp.get('root-log-dir') or hp.get('root') + _checkpoint_auto_every_steps = hp.get('experiment_dump_to_train_steps_ratio') or hp.get('experiment-dump-to-train-steps-ratio') or 0 + except Exception: + _root_log_dir = None + _checkpoint_auto_every_steps = 0 + self._checkpoint_auto_every_steps = int(_checkpoint_auto_every_steps or 0) + + # Initialize CheckpointManagerV2 if we have a root dir (fallback to default root) + root_log_dir = _root_log_dir or os.path.join('.', 'root_log_dir') + try: + # Check if a checkpoint manager is already registered in ledger + try: + existing_manager = ledgers.get_checkpoint_manager() + if existing_manager is not None and not isinstance(existing_manager, ledgers.Proxy): + self._checkpoint_manager = existing_manager + logger.info("Using checkpoint manager from ledger") + else: + raise KeyError("No manager in ledger") + except (KeyError, AttributeError): + # Create new manager and register it + self._checkpoint_manager = CheckpointManagerV2(root_log_dir=root_log_dir) + try: + ledgers.register_checkpoint_manager('default', self._checkpoint_manager) + logger.info("Registered new checkpoint manager in ledger") + except Exception: + pass + except Exception: + self._checkpoint_manager = None + + # Early auto-load latest model architecture and weights if checkpoints exist + if self._checkpoint_manager is not None: + try: + # Try to get the latest experiment hash + latest_hash = None + if hasattr(self._checkpoint_manager, 'current_exp_hash') and self._checkpoint_manager.current_exp_hash: + latest_hash = self._checkpoint_manager.current_exp_hash + elif hasattr(self._checkpoint_manager, 'manifest') and self._checkpoint_manager.manifest: + manifest = self._checkpoint_manager.manifest + latest_hash = getattr(manifest, 'latest_hash', None) + + if latest_hash: + # Use checkpoint manager's load_checkpoint to get architecture and weights + checkpoint_data = self._checkpoint_manager.load_checkpoint( + exp_hash=latest_hash, + load_model=True, + load_weights=True, + load_config=False, + load_data=False, + force=True + ) + + # Apply loaded model if architecture was loaded + if checkpoint_data.get('model'): + self = checkpoint_data['model'] + weights = checkpoint_data.get('weights') + checkpoint_rng_state = checkpoint_data.get('weights', {}).get('rng_state') + + # Restore RNG state if available + restore_rng_state(checkpoint_rng_state) + logger.debug(f"Restored RNG state from checkpoint") + + elif checkpoint_data.get('weights'): + # Only weights available, load into existing model + weights = checkpoint_data['weights'] + if 'model_state_dict' in weights: + self.load_state_dict(weights['model_state_dict'], strict=True) + self.current_step = weights.get('step', -1) + logger.info(f"Auto-loaded model weights from checkpoint {latest_hash[:16]} (step {self.current_step})") + + # As model architecture as has been loaded, and it's an instance of the ModelInterface, + # we can set its current step if available in weights + if isinstance(self.model, self.__class__): + self._registration( + model=self.model, + name=name, + weak=weak + ) + return + + except Exception as e: + logger.debug(f"Could not auto-load model checkpoint: {e}") + if not use_onnx: self.print_graph = print_graph self.print_graph_filename = print_graph_filename - self.traced_model = symbolic_trace(model) + self.traced_model = symbolic_trace(self.model) self.traced_model.name = "N.A." self.guard_training_context = guard_training_context self.guard_testing_context = guard_testing_context @@ -104,26 +212,12 @@ def __init__( # Clean # Optionally register wrapper in global ledger if register: - try: - # Prefer an explicit name. Otherwise prefer a meaningful - # candidate (function __name__ when informative, then - # the class name). Avoid using the generic literal - # 'model' which can be produced by wrappers/patching and - # lead to duplicate registrations. - if name: - reg_name = name - else: - candidate = getattr(model, '__name__', None) - if candidate and candidate.lower() != 'model': - reg_name = candidate - else: - clsname = getattr(model.__class__, '__name__', None) - reg_name = clsname if clsname and clsname.lower() != 'model' else (name or 'model') - - register_model(reg_name, self, weak=weak) - self._ledger_name = reg_name - except Exception: - pass + self._registration( + model=self.model, + name=name, + weak=weak + ) + if not use_onnx: del self.traced_model @@ -136,62 +230,48 @@ def __init__( self.guard_training_context.model = self self.guard_testing_context.model = self - # Checkpoint manager (optional) - # skip_checkpoint_load: bool = False, - # auto_dump_every_steps: int = 0 - self._checkpoint_manager = None - # self._checkpoint_auto_every_steps = int(auto_dump_every_steps or 0) - _checkpoint_auto_every_steps = 0 - _checkpoint_dir = None - _skip_checkpoint_load = False - # If checkpoint_dir not provided, try to read `root_log_dir` from - # ledger hyperparams, otherwise fallback to './root_log_dir/checkpoints' + def _registration(self, model, name, weak: bool = False): try: - from weightslab.backend.ledgers import list_hyperparams, get_hyperparams - names = list_hyperparams() - chosen = None - if 'main' in names: - chosen = 'main' - elif 'experiment' in names: - chosen = 'experiment' - elif len(names) == 1: - chosen = names[0] + # Prefer an explicit name. Otherwise prefer a meaningful + # candidate (function __name__ when informative, then + # the class name). Avoid using the generic literal + # 'model' which can be produced by wrappers/patching and + # lead to duplicate registrations. + if name: + reg_name = name + else: + candidate = getattr(model, '__name__', None) + if candidate and candidate.lower() != 'model': + reg_name = candidate + else: + clsname = getattr(model.__class__, '__name__', None) + reg_name = clsname if clsname and clsname.lower() != 'model' else (name or 'model') - if chosen: - hp = get_hyperparams(chosen) - if hasattr(hp, 'get') and not isinstance(hp, dict): - try: - hp = hp.get() - except Exception: - hp = None - if isinstance(hp, dict): - # Root dir for checkpoints - root = hp.get('root_log_dir') or hp.get('root-log-dir') or hp.get('root') - _checkpoint_dir = os.path.join(str(root), 'checkpoints') if root else None - # Auto dump every N steps - _checkpoint_auto_every_steps = hp.get('experiment_dump_to_train_steps_ratio') or hp.get('experiment-dump-to-train-steps-ratio') or 0 - # Skip loading at init - _skip_checkpoint_load = hp.get('skip_checkpoint_load') or hp.get('skip-checkpoint-load') or False + register_model(reg_name, self, weak=weak) + self._ledger_name = reg_name except Exception: - _checkpoint_dir = None - _checkpoint_auto_every_steps = 0 - _skip_checkpoint_load = False - self._checkpoint_auto_every_steps = int(_checkpoint_auto_every_steps or 0) + pass - if _checkpoint_dir: - try: - self._checkpoint_manager = CheckpointManager(_checkpoint_dir) - # attempt to load latest checkpoint unless skipped - if not _skip_checkpoint_load: - try: - latest = self._checkpoint_manager.get_latest_checkpoint_path() - if latest: - # best-effort load into ledger-registered objects - self._checkpoint_manager.load(str(latest), model_name=(getattr(self, '_ledger_name', None))) - except Exception: - pass - except Exception: - self._checkpoint_manager = None + def load_state_dict(self, state_dict, strict: bool = True): + """ + Loads the state dictionary into the wrapped model. + + This method forwards the provided `state_dict` to the underlying + model's `load_state_dict` method, allowing for the restoration + of model parameters and buffers from a saved state. + + Args: + state_dict (dict): A state dictionary containing model parameters + and buffers to be loaded. + strict (bool, optional): Whether to strictly enforce that the keys + in `state_dict` match the keys returned by the model's + `state_dict()` function. Defaults to True. + + Returns: + None: This method does not return any value; it modifies the + state of the wrapped model in-place. + """ + super().load_state_dict(state_dict, strict=strict) def init_attributes(self, obj): """Expose attributes and methods from the wrapped `obj`. @@ -309,14 +389,34 @@ def _update_optimizer(self, model): def _maybe_auto_dump(self): # Called from base class hook after seen_samples updates. + # Auto-dump: save model weights only (and architecture if changed). try: if not self.is_training() or self._checkpoint_manager is None or self._checkpoint_auto_every_steps <= 0: return batched_age = int(self.get_batched_age()) if batched_age > 0 and (batched_age % self._checkpoint_auto_every_steps) == 0: try: - # best-effort managed dump using ledger names - self._checkpoint_manager.dump(model_name=getattr(self, '_ledger_name', None)) + # Update hash for current experiment state (marks changes as pending, doesn't dump) + new_hash, is_new, changed_components = self._checkpoint_manager.update_experiment_hash() + # If model architecture changed, save it + if 'model' in changed_components: + try: + self._checkpoint_manager.save_model_architecture(self.model) + except Exception: + pass + except Exception: + pass + try: + # Save model weights checkpoint (no pending dump here) + self._checkpoint_manager.save_model_checkpoint( + model=self.model, + model_name=getattr(self, '_ledger_name', None), + save_optimizer=True, + optimizer_name=getattr(self, '_ledger_name', None), + step=batched_age, + force_dump_pending=False, + update_manifest=False + ) except Exception: pass except Exception: @@ -460,7 +560,7 @@ def define_deps(self, use_onnx: bool = False, dummy_input: th.Tensor = None): # Generate the graph dependencies if not use_onnx: - self.shape_propagation() + # self.shape_propagation() self.dependencies_with_ops = generate_graph_dependencies_from_torchfx( self.model, self.traced_model.graph diff --git a/weightslab/components/__init__.py b/weightslab/components/__init__.py index e69de29b..31d64da5 100644 --- a/weightslab/components/__init__.py +++ b/weightslab/components/__init__.py @@ -0,0 +1,51 @@ +""" +Weightslab Components Module + +This module contains core components for experiment tracking, checkpointing, +and monitoring in Weightslab. +""" + +# Legacy checkpoint manager (deprecated) +from weightslab.components.checkpoint import CheckpointManager + +# New structured checkpoint system +from weightslab.components.checkpoint_manager_v2 import CheckpointManagerV2 +from weightslab.components.experiment_hash import ExperimentHashGenerator + +# Automatic checkpoint system (recommended) +from weightslab.components.auto_checkpoint import ( + AutomaticCheckpointSystem, + get_checkpoint_system, + checkpoint_on_step, + checkpoint_on_model_change, + checkpoint_on_config_change, + checkpoint_on_data_change, + checkpoint_on_state_change, +) + +# Other components +from weightslab.components.tracking import Tracker, TrackingMode +# from weightslab.components.global_monitoring import GlobalMonitoring # TODO: Fix missing GlobalMonitoring class + +__all__ = [ + # Checkpoint management + 'CheckpointManager', # Legacy - deprecated + 'CheckpointManagerV2', # Manual checkpoint system + 'ExperimentHashGenerator', + + # Automatic checkpoint system (recommended) + 'AutomaticCheckpointSystem', + 'get_checkpoint_system', + 'checkpoint_on_step', + 'checkpoint_on_model_change', + 'checkpoint_on_config_change', + 'checkpoint_on_data_change', + 'checkpoint_on_state_change', + + # Tracking + 'Tracker', + 'TrackingMode', + + # Monitoring - commented out until GlobalMonitoring is implemented + # 'GlobalMonitoring', +] diff --git a/weightslab/components/auto_checkpoint.py b/weightslab/components/auto_checkpoint.py new file mode 100644 index 00000000..22f5617c --- /dev/null +++ b/weightslab/components/auto_checkpoint.py @@ -0,0 +1,524 @@ +""" +Automatic Checkpoint System - Ledger-Integrated + +This module provides a fully automatic checkpoint management system that: +1. Registers itself in the ledger +2. Auto-initializes on first model/dataloader registration +3. Automatically saves checkpoints every N steps +4. Detects and responds to: + - Model architecture changes → triggers new hash (pending until resume) + - Hyperparameter updates → triggers new hash (pending until resume) + - Data changes (discard, tags) → triggers new hash (pending until resume) + - Model state changes (freeze/reset) → saves metadata + +Changes are marked as "pending" until training resumes or manual dump is forced. + +The system is completely transparent to the user - no manual calls needed. +""" + +import logging +import threading +from pathlib import Path +from typing import Any, Dict, Optional, Set +from datetime import datetime + +import torch as th + +from weightslab.components.checkpoint_manager_v2 import CheckpointManagerV2 +from weightslab.components.experiment_hash import ExperimentHashGenerator +from weightslab.backend import ledgers + + +logger = logging.getLogger(__name__) + + +class AutomaticCheckpointSystem: + """Automatic checkpoint system that integrates with the ledger. + + This system: + - Monitors model, optimizer, and hyperparameter registrations + - Automatically saves checkpoints every N steps + - Detects configuration changes and creates new checkpoint directories + - Completely transparent to the user + + Attributes: + checkpoint_manager (CheckpointManagerV2): Core checkpoint manager + checkpoint_frequency (int): Save checkpoints every N steps + _initialized (bool): Whether system has been initialized + _step_counter (int): Current training step + _lock (threading.Lock): Thread safety lock + """ + + def __init__( + self, + root_log_dir: str = 'root_experiment', + checkpoint_frequency: int = 100, + auto_register: bool = True + ): + """Initialize the automatic checkpoint system. + + Args: + root_log_dir: Root directory for experiments + checkpoint_frequency: Save checkpoints every N steps + auto_register: Auto-register in ledger on init + """ + self.checkpoint_manager = CheckpointManagerV2(root_log_dir=root_log_dir) + self.checkpoint_frequency = checkpoint_frequency + + self._initialized = False + self._step_counter = 0 + self._lock = threading.Lock() + + # Track last known states for change detection + self._last_model_id = None + self._last_config = None + self._last_data_state = None + self._last_checkpoint_step = -1 + self._training_resumed = False + + logger.info(f"AutomaticCheckpointSystem initialized (freq={checkpoint_frequency})") + + def initialize_from_ledger(self): + """Initialize checkpoint system from current ledger state. + + This is called automatically on first checkpoint or can be called + manually to sync with ledger. + """ + with self._lock: + if self._initialized: + logger.debug("Checkpoint system already initialized") + return + + try: + model = self._get_model_from_ledger() + config = self._get_config_from_ledger() + data = self._get_dfm_from_ledger() + if model is not None or config is not None: + exp_hash, is_new, changed = self.checkpoint_manager.update_experiment_hash( + model_snapshot=model, + hp_snapshot=config, + dfm_snapshot=data, + ) + + self._last_config = config.copy() if config else None + self._initialized = True + + logger.info(f"Checkpoint system initialized with hash: {exp_hash}") + else: + logger.warning("No model or config found in ledger for initialization") + + except Exception as e: + logger.error(f"Failed to initialize from ledger: {e}") + + def on_training_step(self, step: Optional[int] = None, force_dump: bool = False): + """Called after each training step to potentially save checkpoint. + + This: + 1. Dumps pending changes if training resumed + 2. Checks if it's time for a periodic checkpoint + 3. Can force dump pending changes if requested + + Args: + step: Training step number (auto-increments if None) + force_dump: Force dump pending changes regardless of frequency + """ + with self._lock: + if not self._initialized: + self.initialize_from_ledger() + + if step is not None: + self._step_counter = step + else: + self._step_counter += 1 + + current_step = self._step_counter + + if not self._training_resumed: + has_pending, pending_comps = self.checkpoint_manager.has_pending_changes() + if has_pending: + logger.info(f"Training resumed, dumping pending changes: {pending_comps}") + self.checkpoint_manager.dump_pending_changes(force=True) + self._training_resumed = True + + if force_dump: + self.checkpoint_manager.dump_pending_changes(force=True) + + if current_step % self.checkpoint_frequency == 0: + if current_step != self._last_checkpoint_step: + self._save_checkpoint(step=current_step, force_dump_pending=force_dump) + self._last_checkpoint_step = current_step + + def on_model_change(self, model: Optional[th.nn.Module] = None, dump_immediately: bool = False): + """Called when model architecture changes (add/prune layers). + + By default, changes are marked as pending until training resumes. + Set dump_immediately=True to dump right away. + + Args: + model: New model (gets from ledger if None) + dump_immediately: If True, dump immediately instead of marking pending + """ + with self._lock: + if model is None: + model = self._get_model_from_ledger() + + if model is None: + return + + config = self._get_config_from_ledger() + data_state = self._get_data_state_from_ledger() + + exp_hash, is_new, changed_components = self.checkpoint_manager.update_experiment_hash( + model_snapshot=model, + hp_snapshot=config, + dfm_snapshot=data_state, + dump_immediately=dump_immediately + ) + + if is_new: + logger.info(f"Model architecture changed, new hash: {exp_hash}") + if dump_immediately: + logger.info("Changes dumped immediately") + else: + logger.info("Changes marked as pending (will dump on training resume)") + + self._training_resumed = False + + def on_config_change(self, config: Optional[Dict[str, Any]] = None, dump_immediately: bool = False): + """Called when hyperparameters change. + + By default, changes are marked as pending until training resumes. + Set dump_immediately=True to dump right away. + + Args: + config: New config (gets from ledger if None) + dump_immediately: If True, dump immediately instead of marking pending + """ + with self._lock: + if config is None: + config = self._get_config_from_ledger() + + if config is None: + return + + if self._last_config is not None and config == self._last_config: + logger.debug("Config unchanged, skipping checkpoint") + return + + model = self._get_model_from_ledger() + data_state = self._get_data_state_from_ledger() + + exp_hash, is_new, changed_components = self.checkpoint_manager.update_experiment_hash( + model_snapshot=model, + hp_snapshot=config, + dfm_snapshot=data_state, + dump_immediately=dump_immediately + ) + + if is_new: + logger.info(f"Hyperparameters changed, new hash: {exp_hash}") + if dump_immediately: + logger.info("Changes dumped immediately") + else: + logger.info("Changes marked as pending (will dump on training resume)") + + self._last_config = config.copy() if config else None + self._training_resumed = False + + def on_data_change(self, data_state: Optional[Dict[str, Any]] = None, dump_immediately: bool = False): + """Called when data state changes (discard, tags). + + By default, changes are marked as pending until training resumes. + Set dump_immediately=True to dump right away. + + Args: + data_state: Dict with 'uids', 'discarded', 'tags' (gets from ledger if None) + dump_immediately: If True, dump immediately instead of marking pending + """ + with self._lock: + if data_state is None: + data_state = self._get_data_state_from_ledger() + + if data_state is None: + return + + if self._last_data_state is not None and data_state == self._last_data_state: + logger.debug("Data state unchanged, skipping checkpoint") + return + + model = self._get_model_from_ledger() + config = self._get_config_from_ledger() + + exp_hash, is_new, changed_components = self.checkpoint_manager.update_experiment_hash( + model_snapshot=model, + hp_snapshot=config, + dfm_snapshot=data_state, + dump_immediately=dump_immediately + ) + + if is_new: + logger.info(f"Data state changed, new hash: {exp_hash}") + if dump_immediately: + logger.info("Changes dumped immediately") + else: + logger.info("Changes marked as pending (will dump on training resume)") + + self._last_data_state = data_state if data_state else None + self._training_resumed = False + + def on_change(self, dump_immediately: bool = False): + """Called when any tracked component changes. + + This is a convenience method that checks model, config, and data. + """ + self.on_model_change(dump_immediately=dump_immediately) + self.on_config_change(dump_immediately=dump_immediately) + self.on_data_change(dump_immediately=dump_immediately) + + def on_model_state_change(self, event_type: str): + """Called when model state changes (freeze, reset, etc.). + + This triggers a checkpoint save with metadata. + + Args: + event_type: Type of state change ('freeze', 'reset', etc.) + """ + with self._lock: + logger.info(f"Model state change: {event_type}") + + model = self._get_model_from_ledger() + if model is not None: + self.checkpoint_manager.save_model_checkpoint( + model=model, + step=self._step_counter, + save_optimizer=True, + metadata={ + 'trigger': 'state_change', + 'event_type': event_type, + 'timestamp': datetime.now().isoformat() + } + ) + + def _save_checkpoint(self, step: int, force_dump_pending: bool = False): + """Internal method to save a checkpoint. + + Args: + step: Training step number + force_dump_pending: Force dump pending changes before saving + """ + try: + model = self._get_model_from_ledger() + + if model is None: + logger.warning("No model found in ledger, skipping checkpoint") + return + + checkpoint_path = self.checkpoint_manager.save_model_checkpoint( + model=model, + step=step, + save_optimizer=True, + metadata={'step': step}, + force_dump_pending=force_dump_pending + ) + + if checkpoint_path: + logger.info(f"Saved checkpoint at step {step}") + else: + logger.warning(f"Failed to save checkpoint at step {step}") + + except Exception as e: + logger.error(f"Error saving checkpoint: {e}") + + def _get_dfm_from_ledger(self) -> Optional[th.nn.Module]: + """Get DataFrame Manager from ledger, handling proxies. + Returns: + Dataframe Manager + """ + try: + dfm = ledgers.get_dataframe() + + if dfm is not None: + return dfm + return None + + except Exception as e: + logger.debug(f"Could not get dataframe manager from ledger: {e}") + return None + + def _get_model_from_ledger(self) -> Optional[th.nn.Module]: + """Get model from ledger, handling proxies. + + Returns: + PyTorch model or None + """ + try: + model = ledgers.get_model() + + if model is not None: + return model + + return None + + except Exception as e: + logger.debug(f"Could not get model from ledger: {e}") + return None + + def _get_config_from_ledger(self) -> Optional[Dict[str, Any]]: + """Get hyperparameters from ledger, handling proxies. + + Returns: + Config dict or None + """ + try: + config = ledgers.get_hyperparams() + + if config is not None: + return config + + return None + + except Exception as e: + logger.debug(f"Could not get config from ledger: {e}") + return None + + def _get_data_state_from_ledger(self) -> Optional[Dict[str, Any]]: + """Get data state from dataloaders in ledger. + + Aggregates UIDs, discard status, and tags from all registered dataloaders. + + Returns: + Dict with 'uids', 'discarded', 'tags' or None + """ + try: + dfm = ledgers.get_dataframe() + + if dfm is not None: + return dfm + + return None + + except Exception as e: + logger.debug(f"Could not get dfm from ledger: {e}") + return None + + def get_status(self) -> Dict[str, Any]: + """Get current status of the checkpoint system. + + Returns: + dict: Status information + """ + with self._lock: + return { + 'initialized': self._initialized, + 'current_step': self._step_counter, + 'last_checkpoint_step': self._last_checkpoint_step, + 'checkpoint_frequency': self.checkpoint_frequency, + 'current_exp_hash': self.checkpoint_manager.current_exp_hash, + 'root_log_dir': str(self.checkpoint_manager.root_log_dir) + } + + +_GLOBAL_CHECKPOINT_SYSTEM: Optional[AutomaticCheckpointSystem] = None +_SYSTEM_LOCK = threading.Lock() + + +def get_checkpoint_system( + root_log_dir: Optional[str] = None, + checkpoint_frequency: int = 100, + auto_init: bool = True +) -> AutomaticCheckpointSystem: + """Get or create the global automatic checkpoint system. + + Args: + root_log_dir: Root directory (only used on first call) + checkpoint_frequency: Checkpoint frequency (only used on first call) + auto_init: Auto-initialize from ledger + + Returns: + AutomaticCheckpointSystem: Global checkpoint system instance + """ + global _GLOBAL_CHECKPOINT_SYSTEM + + with _SYSTEM_LOCK: + if _GLOBAL_CHECKPOINT_SYSTEM is None: + if root_log_dir is None: + try: + hp = ledgers.get_hyperparams() + root_log_dir = hp.get('root_log_dir', 'root_experiment') if hp else 'root_experiment' + except Exception: + root_log_dir = 'root_experiment' + + _GLOBAL_CHECKPOINT_SYSTEM = AutomaticCheckpointSystem( + root_log_dir=root_log_dir, + checkpoint_frequency=checkpoint_frequency, + auto_register=True + ) + + if auto_init: + _GLOBAL_CHECKPOINT_SYSTEM.initialize_from_ledger() + + return _GLOBAL_CHECKPOINT_SYSTEM + + +def checkpoint_on_step(step: Optional[int] = None, force_dump: bool = False): + """Convenience function to trigger checkpoint on training step. + + This can be called from training loops. Use force_dump=True to + immediately dump any pending changes. + + Args: + step: Training step number + force_dump: Force dump pending changes + """ + system = get_checkpoint_system() + system.on_training_step(step=step, force_dump=force_dump) + + +def checkpoint_on_model_change(model: Optional[th.nn.Module] = None, dump_immediately: bool = False): + """Convenience function to trigger checkpoint on model architecture change. + + Args: + model: New model (gets from ledger if None) + dump_immediately: Dump changes immediately instead of marking pending + """ + system = get_checkpoint_system() + system.on_model_change(model=model, dump_immediately=dump_immediately) + + +def checkpoint_on_config_change(config: Optional[Dict[str, Any]] = None, dump_immediately: bool = False): + """Convenience function to trigger checkpoint on config change. + + Args: + config: New config (gets from ledger if None) + dump_immediately: Dump changes immediately instead of marking pending + """ + system = get_checkpoint_system() + system.on_config_change(config=config, dump_immediately=dump_immediately) + + +def checkpoint_on_data_change(data_state: Optional[Dict[str, Any]] = None, dump_immediately: bool = False): + """Convenience function to trigger checkpoint on data state change. + + Args: + data_state: Dict with 'uids', 'discarded', 'tags' (gets from ledger if None) + dump_immediately: Dump changes immediately instead of marking pending + """ + system = get_checkpoint_system() + system.on_data_change(data_state=data_state, dump_immediately=dump_immediately) + + +def checkpoint_on_state_change(event_type: str): + """Convenience function to trigger checkpoint on model state change. + + Args: + event_type: Type of state change ('freeze', 'reset', etc.) + """ + system = get_checkpoint_system() + system.on_model_state_change(event_type=event_type) + +def checkpoint_on_change(dump_immediately: bool = False): + """Convenience function to trigger checkpoint on any tracked component change. + """ + logger.info('\nCheck if changes to dump.') + system = get_checkpoint_system(auto_init=False) + system.on_change(dump_immediately=dump_immediately) diff --git a/weightslab/components/checkpoint_manager_v2.py b/weightslab/components/checkpoint_manager_v2.py new file mode 100644 index 00000000..5f010e6c --- /dev/null +++ b/weightslab/components/checkpoint_manager_v2.py @@ -0,0 +1,1623 @@ +""" +Checkpoint Manager V2 - Structured Checkpoint Management + + +This module implements checkpoint management with separated component directories: + +Directory Structure: + root_log_dir/ + data/ # Data-related files (global) + logs/ # Training logs (global) + checkpoints/ + manifest.yaml # Tracks all hashes with timestamps + models/ + {hash}/ # 24-byte hash: HP_MODEL_DATA + {hash}_step_000100.pt + {hash}_architecture.pkl + HP/ + {hash}/ + {hash}_config.yaml + data/ + {hash}/ + {hash}_data_state.yaml + +Hash format: HP(8) + MODEL(8) + DATA(8) = 24 bytes + +Manifest tracks hash chronology for loading most recent experiments. +""" + +import os +import json +import yaml +import logging +import shutil +import random +import numpy as np +from pathlib import Path +from typing import Any, Dict, List, Optional, Set +from datetime import datetime +import pandas as pd + +import torch as th +import dill +import pickle + +from weightslab.components.global_monitoring import guard_training_context, guard_testing_context +from weightslab.components.experiment_hash import ExperimentHashGenerator +from weightslab.backend.ledgers import ( + get_model, + get_optimizer, + get_dataloader, + get_dataloaders, +) +from weightslab.backend import ledgers +from weightslab.data.sample_stats import SampleStatsEx +from weightslab.utils.tools import capture_rng_state, restore_rng_state +from weightslab.components.global_monitoring import pause_controller as pause_ctrl + +# Init logger +logger = logging.getLogger(__name__) + + +class CheckpointManagerV2: + """Structured checkpoint manager with experiment hash-based organization. + + This manager creates a well-organized checkpoint structure where each + unique experiment configuration gets its own directory identified by + a deterministic hash. + + Attributes: + root_log_dir (Path): Root directory for all experiment outputs + checkpoints_dir (Path): Base checkpoints directory + hash_generator (ExperimentHashGenerator): Hash generation utility + current_exp_hash (str): Current experiment hash + _step_counter (int): Global step counter for model checkpoints + """ + + def __init__(self, root_log_dir: str = 'root_experiment'): + """Initialize the checkpoint manager. + + Args: + root_log_dir: Root directory for experiment outputs + """ + self.root_log_dir = Path(root_log_dir).absolute() + self.root_log_dir.mkdir(parents=True, exist_ok=True) + + # Create main subdirectories + self.data_dir = self.root_log_dir / "data" + self.logs_dir = self.root_log_dir / "logs" + self.checkpoints_dir = self.root_log_dir / "checkpoints" + + self.data_dir.mkdir(exist_ok=True) + self.logs_dir.mkdir(exist_ok=True) + self.checkpoints_dir.mkdir(exist_ok=True) + + # Create checkpoint subdirectories for different components + self.models_dir = self.checkpoints_dir / "models" + self.hp_dir = self.checkpoints_dir / "HP" + self.data_checkpoint_dir = self.checkpoints_dir / "data" + self.loggers_dir = self.checkpoints_dir / "loggers" + + self.models_dir.mkdir(exist_ok=True) + self.hp_dir.mkdir(exist_ok=True) + self.data_checkpoint_dir.mkdir(exist_ok=True) + self.loggers_dir.mkdir(exist_ok=True) + + # Manifest file for tracking hash chronology + self.manifest_file = self.checkpoints_dir / "manifest.yaml" + + # Hash management + self.hash_generator = ExperimentHashGenerator() + self.current_exp_hash: Optional[str] = None + self.previous_exp_hash: Optional[str] = None + self.hash_by_module: list = [None, None, None] # HP, MODEL, DATA + + # Step tracking + self._step_counter = 0 + + # First time only + self.firsttime = True + + # Pending changes tracking + self._pending_model = None + self._pending_config = None + self._pending_data_state = None + self._has_pending_changes = False + self._pending_components = set() + + # Load existing state if available + self._load_manager_state() + + # Load any existing logger snapshots for visibility when starting + self._load_all_logger_snapshots() + + # Automatically resume latest state when an existing root_log_dir is provided + self._bootstrap_latest_state() + + logger.info(f"CheckpointManagerV2 initialized at {self.root_log_dir}") + + def __repr__(self) -> str: + return ( + f"CheckpointManagerV2(\n" + f" root_log_dir={self.root_log_dir}\n" + f" current_exp_hash={self.current_exp_hash}\n" + f" step_counter={self._step_counter}\n" + f")" + ) + + def _get_data_state_snapshot(self, dfm): + """Return a combined data state from registered dataloaders, if present.""" + try: + if isinstance(dfm, dict): + return dfm + + collected_discarded = {} + collected_tags = {} + + collected_discarded.update(dfm.get_df_view(SampleStatsEx.DENY_LISTED.value).to_dict()) + collected_tags.update(dfm.get_df_view(SampleStatsEx.TAGS.value).to_dict()) + + if not collected_tags and not collected_discarded: + return None + + return { + 'discarded': collected_discarded, + 'tags': collected_tags, + } + except Exception: + return None + + def get_current_experiment_hash(self) -> Optional[str]: + """Get the current experiment hash.""" + return self.current_exp_hash + + def get_HP_snapshot(self) -> Dict[str, Any]: + """Get current hyperparameters snapshot from ledger.""" + try: + hp = ledgers.get_hyperparams() + if hp is None: + return {} + if isinstance(hp, ledgers.Proxy) and hasattr(hp, 'get') and callable(hp.get): + hp = hp.get() + if isinstance(hp, dict): + return hp + elif hasattr(hp, '__dict__'): + return vars(hp) + else: + return {} + except Exception: + return {} + + def get_model_snapshot(self) -> Optional[th.nn.Module]: + """Get current model snapshot from ledger.""" + try: + model = ledgers.get_model() + if model is None: + return None + if isinstance(model, ledgers.Proxy) and hasattr(model, 'get') and callable(model.get): + model = model.get() + if isinstance(model, th.nn.Module): + return model + else: + return None + except Exception: + return None + + def get_dataframe_snapshot(self) -> Optional[Dict[str, Any]]: + """Get current dataframe snapshot from registered dataloaders.""" + try: + dfm = ledgers.get_dataframe() + if isinstance(dfm, ledgers.Proxy) and hasattr(dfm, 'get') and callable(dfm.get): + dfm = dfm.get() + if dfm is None: + return None + return self._get_data_state_snapshot(dfm) + except Exception: + return None + + def update_experiment_hash( + self, + model_snapshot: Optional[th.nn.Module] = None, + hp_snapshot: Optional[Dict[str, Any]] = None, + dfm_snapshot: Optional[Dict[str, Any]] = None, + force: bool = False, + firsttime: bool = False, + dump_immediately: bool = False + ) -> tuple[str, bool, Set[str]]: + """Update experiment hash and track changes (pending or immediate). + + Changes can be: + 1. Pending: Tracked but not dumped until training resumes or manual dump + 2. Immediate: Dumped right away if dump_immediately=True + + Args: + model_snapshot: PyTorch model + hp_snapshot: Dictionary of hyperparameters + dfm_snapshot: Dictionary with 'uids', 'discarded', 'tags' + force: Force hash regeneration even if nothing changed + dump_immediately: If True, dump changes immediately. If False, mark as pending. + + Returns: + tuple: (exp_hash: str, is_new: bool, changed_components: Set[str]) + """ + # Init first time saving the init state when first resumes + if firsttime and self.firsttime and not force: + logger.info("First time initialization; skipping hash update.") + self.firsttime = False + force = True + dump_immediately = True + + # Get ledgered components + hp_snapshot = self.get_HP_snapshot() if hp_snapshot is None else hp_snapshot + data_snapshot = self.get_dataframe_snapshot() if dfm_snapshot is None else dfm_snapshot + model_snapshot = self.get_model_snapshot() if model_snapshot is None else model_snapshot + + # Check what changed + has_changed, changed_components = self.hash_generator.has_changed( + model=model_snapshot, + config=hp_snapshot, + data_state=data_snapshot, + force=force + ) + + if not has_changed and not force: + return self.current_exp_hash, False, set() + + # Generate new hash with all components + new_hash = self.hash_generator.generate_hash( + model=model_snapshot, + config=hp_snapshot, + data_state=data_snapshot + ) + + is_new = (new_hash != self.current_exp_hash) or (force or dump_immediately) + + if is_new: + logger.info(f"New experiment hash: {new_hash} (previous: {self.current_exp_hash})") + logger.info(f"Changed components: {changed_components}") + + # Update hash + old_hash = self.current_exp_hash + self.current_exp_hash = new_hash + self.previous_exp_hash = old_hash + self.hash_by_module[0] = self.hash_generator.get_component_hashes().get('hp', None) + self.hash_by_module[1] = self.hash_generator.get_component_hashes().get('model', None) + self.hash_by_module[2] = self.hash_generator.get_component_hashes().get('data', None) + + if dump_immediately: + # Dump changes immediately + self._dump_changes( + model=model_snapshot, + config=hp_snapshot, + data_state=data_snapshot, + changed_components=changed_components + ) + self._has_pending_changes = False + self._pending_components = set() + + # Sync RNG and data state from just-dumped checkpoint + try: + # Restore RNG state + rng_state = capture_rng_state() + restore_rng_state(rng_state) + + # Reset dataloader iterators to sync with new state + for loader_name in get_dataloaders(): + loader = get_dataloader(loader_name) + if hasattr(loader, 'reset_iterator') and callable(loader.reset_iterator): + loader.reset_iterator() + logger.debug(f"Reset iterator for dataloader: {loader_name}") + except Exception as e: + logger.warning(f"Failed to sync RNG/data state after dump: {e}") + else: + # Mark as pending + self._pending_model = model_snapshot + self._pending_config = hp_snapshot + self._pending_data_state = data_snapshot + self._has_pending_changes = True + self._pending_components = changed_components + logger.info(f"Changes pending (not dumped yet): {changed_components}") + + # Save manager state + self._save_manager_state() + + return new_hash, is_new, changed_components + + def _create_exp_hash_directories( + self, + exp_hash: str, + create_model_dir: bool = True, + create_hp_dir: bool = True, + create_data_dir: bool = True + ): + """Create directory structure for an experiment hash in separate component folders. + + Args: + exp_hash: 24-byte experiment hash (HP_MODEL_DATA) + """ + model_hash_dir = self.models_dir / exp_hash + hp_hash_dir = self.hp_dir / exp_hash + data_hash_dir = self.data_checkpoint_dir / exp_hash + + if create_model_dir: + model_hash_dir.mkdir(exist_ok=True) + if create_hp_dir: + hp_hash_dir.mkdir(exist_ok=True) + if create_data_dir: + data_hash_dir.mkdir(exist_ok=True) + + logger.debug(f"Created checkpoint directories for {exp_hash}") + self._update_manifest(exp_hash) + + def _update_manifest(self, exp_hash: str): + """Update manifest file with new or updated hash.""" + try: + manifest = self._load_manifest() + component_hashes = self.hash_generator.get_component_hashes() + + if exp_hash not in manifest['experiments']: + manifest['experiments'][exp_hash] = { + 'hp_hash': component_hashes.get('hp', exp_hash[0:8]), + 'model_hash': component_hashes.get('model', exp_hash[8:16]), + 'data_hash': component_hashes.get('data', exp_hash[16:24]), + 'created': datetime.now().isoformat(), + 'last_used': datetime.now().isoformat(), + 'latest_weight_checkpoint': None, + 'latest_weight_step': None + } + else: + manifest['experiments'][exp_hash]['last_used'] = datetime.now().isoformat() + + manifest['latest_hash'] = exp_hash + manifest['last_updated'] = datetime.now().isoformat() + + # Ensure parent directory exists + self.manifest_file.parent.mkdir(parents=True, exist_ok=True) + + # Write manifest file (overwrites if exists) + with open(self.manifest_file, 'w') as f: + yaml.dump(manifest, f, default_flow_style=False, sort_keys=False) + except Exception as e: + logger.warning(f"Failed to update manifest: {e}") + + def _update_manifest_weight_checkpoint(self, exp_hash: str, checkpoint_filename: str, step: int): + """Update manifest with latest weight checkpoint for given experiment hash.""" + try: + manifest = self._load_manifest() + if exp_hash in manifest['experiments']: + manifest['experiments'][exp_hash]['latest_weight_checkpoint'] = checkpoint_filename + manifest['experiments'][exp_hash]['latest_weight_step'] = step + manifest['experiments'][exp_hash]['last_used'] = datetime.now().isoformat() + + # Write updated manifest + with open(self.manifest_file, 'w') as f: + yaml.dump(manifest, f, default_flow_style=False, sort_keys=False) + logger.debug(f"Updated manifest with weight checkpoint: {checkpoint_filename} (step {step})") + except Exception as e: + logger.warning(f"Failed to update manifest weight checkpoint: {e}") + + # ------------------------------------------------------------------ + # Logger snapshot management + # ------------------------------------------------------------------ + def _get_logger_snapshot_path(self, exp_hash: Optional[str] = None) -> Path: + exp = exp_hash or self.current_exp_hash + return self.loggers_dir / exp / "loggers.json" if exp else None + + def save_logger_snapshot(self, exp_hash: Optional[str] = None) -> Optional[Path]: + """Persist logger queues for the given experiment hash. + + Uses the same hash as model/hp/data; does not affect hashing. + """ + exp = exp_hash or self.current_exp_hash + if exp is None: + return None + + try: + logger_names = ledgers.list_loggers() + if not logger_names: + return None + + snapshot = {"exp_hash": exp, "timestamp": datetime.now().isoformat(), "loggers": {}} + for lname in logger_names: + lg = ledgers.get_logger(lname) + if lg is None: + continue + # Expect LoggerQueue interface + history = lg.get_signal_history() if hasattr(lg, "get_signal_history") else [] + graphs = lg.get_graph_names() if hasattr(lg, "get_graph_names") else [] + snapshot["loggers"][lname] = { + "signal_history": history, + "graph_names": graphs, + } + + if not snapshot["loggers"]: + return None + + path = self._get_logger_snapshot_path(exp) + path.parent.mkdir(parents=True, exist_ok=True) + with open(path, "w") as f: + json.dump(snapshot, f, indent=2) + logger.info(f"Saved logger snapshot: {path}") + return path + except Exception as e: + logger.warning(f"Failed to save logger snapshot: {e}") + return None + + def load_logger_snapshot(self, exp_hash: str) -> bool: + """Load logger queues from snapshot for a specific experiment hash.""" + path = self._get_logger_snapshot_path(exp_hash) + if path is None or not path.exists(): + return False + try: + with open(path, "r") as f: + snapshot = json.load(f) + + loggers_payload = snapshot.get("loggers", {}) + for lname, payload in loggers_payload.items(): + try: + from weightslab.utils.logger import LoggerQueue + + lg = ledgers.get_logger(lname) if lname in ledgers.list_loggers() else None + # If no real logger (or only a proxy) exists, create a fresh LoggerQueue so we can enqueue history + if lg is None or not hasattr(lg, "load_snapshot"): + lg = LoggerQueue(name=lname, register=True) + lg.load_snapshot(payload) + except Exception as inner_e: + logger.warning(f"Failed to restore logger '{lname}': {inner_e}") + return True + except Exception as e: + logger.warning(f"Failed to load logger snapshot for {exp_hash}: {e}") + return False + + def _load_all_logger_snapshots(self): + """Load all logger snapshots found under loggers/ for visibility when starting.""" + if not self.loggers_dir.exists(): + return + for exp_dir in self.loggers_dir.iterdir(): + if exp_dir.is_dir(): + snapshot_file = exp_dir / "loggers.json" + if snapshot_file.exists(): + self.load_logger_snapshot(exp_dir.name) + + def _load_manifest(self) -> Dict[str, Any]: + """Load manifest file.""" + if self.manifest_file.exists(): + try: + with open(self.manifest_file, 'r') as f: + return yaml.safe_load(f) or {'experiments': {}, 'latest_hash': None} + except Exception as e: + logger.warning(f"Failed to load manifest: {e}") + return {'experiments': {}, 'latest_hash': None} + + def get_latest_hash(self) -> Optional[str]: + """Get the most recent experiment hash.""" + manifest = self._load_manifest() + return manifest.get('latest_hash') + + def get_all_hashes(self, sort_by: str = 'created') -> List[Dict[str, Any]]: + """Get all hashes sorted by timestamp.""" + manifest = self._load_manifest() + experiments = manifest.get('experiments', {}) + hash_list = [{'hash': h, **info} for h, info in experiments.items()] + if sort_by in ['created', 'last_used']: + hash_list.sort(key=lambda x: x.get(sort_by, ''), reverse=True) + return hash_list + + def get_hashes_by_component(self, hp_hash: Optional[str] = None, + model_hash: Optional[str] = None, + data_hash: Optional[str] = None) -> List[str]: + """Find hashes matching component hash(es).""" + manifest = self._load_manifest() + experiments = manifest.get('experiments', {}) + matching = [] + for exp_hash, info in experiments.items(): + if hp_hash and info.get('hp_hash') != hp_hash: + continue + if model_hash and info.get('model_hash') != model_hash: + continue + if data_hash and info.get('data_hash') != data_hash: + continue + matching.append(exp_hash) + return matching + + def dump_pending_changes(self, force: bool = False) -> bool: + """Dump any pending changes to disk. + + This is called when: + - Training resumes after model/config/data changes + - Manual checkpoint is requested with force=True + + Args: + force: Force dump even if no pending changes + + Returns: + bool: True if changes were dumped, False otherwise + """ + if not self._has_pending_changes and not force: + logger.debug("No pending changes to dump") + return False + + if self.current_exp_hash is None: + logger.warning("No experiment hash set. Cannot dump pending changes.") + return False + + logger.info(f"Dumping pending changes: {self._pending_components}") + + self._dump_changes( + model=self._pending_model, + config=self._pending_config, + data_state=self._pending_data_state, + changed_components=self._pending_components + ) + + # Clear pending state + self._pending_model = None + self._pending_config = None + self._pending_data_state = None + self._has_pending_changes = False + self._pending_components = set() + + self._save_manager_state() + return True + + def _dump_changes( + self, + model: Optional[th.nn.Module], + config: Optional[Dict[str, Any]], + data_state: Optional[Dict[str, Any]], + changed_components: Set[str] + ): + """Internal method to dump changes to disk. + + Args: + model: PyTorch model + config: Hyperparameters config + data_state: Data state dict + changed_components: Set of changed components ('model', 'config', 'data') + """ + # Create checkpoint subdirectories for this hash + self._create_exp_hash_directories( + self.current_exp_hash, + create_data_dir='data' in changed_components, + create_hp_dir='config' in changed_components, + create_model_dir='model' in changed_components + ) + + # Track if we need to save weights + should_save_weights = False + weights_model = None + + if 'model' in changed_components and model is not None: + logger.info("Dumping model architecture...") + self.save_model_architecture(model) + should_save_weights = True + weights_model = model + + if ('hp' in changed_components or 'config' in changed_components) and config is not None: + logger.info("Dumping hyperparameters config...") + self.save_config(config) + should_save_weights = True + + if 'data' in changed_components: + logger.info("Dumping data snapshot...") + self.save_data_snapshot() + should_save_weights = True + + # Save weights whenever any component changes to preserve complete state + if should_save_weights: + try: + # Get model from ledger if not provided + if weights_model is None: + try: + weights_model = get_model() + if hasattr(weights_model, 'get') and callable(weights_model.get): + weights_model = weights_model.get() + except Exception: + pass + + if weights_model is not None: + logger.info("Saving model weights checkpoint with component changes...") + self.save_model_checkpoint(weights_model) + else: + logger.warning("Could not save weights: no model available") + except Exception as e: + logger.warning(f"Failed to save weights with component changes: {e}") + + # Always save logger snapshot alongside other components (same hash) + self.save_logger_snapshot() + + def has_pending_changes(self) -> tuple[bool, Set[str]]: + """Check if there are pending changes. + + Returns: + tuple: (has_pending: bool, pending_components: Set[str]) + """ + return self._has_pending_changes, self._pending_components.copy() + + # ================ + # SAVING FUNCTIONS + # ================ + def save_model_checkpoint( + self, + model: Optional[th.nn.Module] = None, + model_name: Optional[str] = None, + step: Optional[int] = None, + save_optimizer: bool = True, + optimizer_name: Optional[str] = None, + metadata: Optional[Dict[str, Any]] = None, + force_dump_pending: bool = False, + update_manifest: bool = True + ) -> Optional[Path]: + """Save model weights checkpoint. + + This saves only the model weights (state_dict) for fast checkpointing + during training. The architecture is saved separately when the hash changes. + + Args: + model: PyTorch model (or get from ledger if None) + model_name: Name to get model from ledger + step: Training step number (uses internal counter if None) + save_optimizer: Whether to also save optimizer state + optimizer_name: Name to get optimizer from ledger + metadata: Additional metadata to save + force_dump_pending: If True, dump any pending changes before saving checkpoint + + Returns: + Path: Path to saved checkpoint file, or None if failed + """ + if self.current_exp_hash is None: + logger.warning("No experiment hash set. Call update_experiment_hash first.") + return None + + # Dump pending changes if requested + if force_dump_pending and self._has_pending_changes: + logger.info("Force dumping pending changes before checkpoint...") + self.dump_pending_changes(force=True) + + # Get model from ledger if not provided + if model is None: + try: + model = get_model(model_name) + # Unwrap proxy if needed + if hasattr(model, 'get') and callable(model.get): + model = model.get() + except Exception as e: + logger.error(f"Could not get model from ledger: {e}") + return None + + if model is None: + logger.error("No model available to checkpoint") + return None + + # Determine step + if step is None: + step = self._step_counter + self._step_counter = max(self._step_counter, step + 1) + + # Prepare checkpoint data + checkpoint = { + 'step': step, + 'model_state_dict': model.state_dict(), + 'timestamp': datetime.now().isoformat(), + 'exp_hash': self.current_exp_hash, + 'rng_state': capture_rng_state(), # Capture RNG state for reproducible training + } + + # Capture dataloader iteration state(s) for reproducible resume (support multiple loaders) + try: + loader_states = {} + for loader_name in get_dataloaders(): + dataloader = get_dataloader(loader_name) + if dataloader is not None and hasattr(dataloader, 'capture_iteration_state'): + loader_states[loader_name] = dataloader.capture_iteration_state() + if loader_states: + checkpoint['dataloader_iteration_state'] = loader_states + logger.debug(f"Captured dataloader iteration states: {loader_states}") + except Exception as e: + logger.debug(f"Could not capture dataloader iteration state: {e}") + + # Add optimizer state if requested + if save_optimizer: + try: + optimizer = get_optimizer(optimizer_name) + if hasattr(optimizer, 'get') and callable(optimizer.get): + optimizer = optimizer.get() + if optimizer is not None: + checkpoint['optimizer_state_dict'] = optimizer.state_dict() + except Exception as e: + logger.warning(f"Could not save optimizer state: {e}") + + # Add metadata + if metadata: + checkpoint['metadata'] = metadata + + # Save checkpoint + model_dir = self.models_dir / self.current_exp_hash[8:-8] + os.makedirs(model_dir, exist_ok=True) + # Use full exp_hash in filename for clarity and uniqueness + checkpoint_file = model_dir / f"{self.current_exp_hash}_step_{step:06d}.pt" + + try: + th.save(checkpoint, checkpoint_file) + logger.info(f"Saved model checkpoint: {checkpoint_file.name}") + + # Update manifest with latest weight checkpoint for this experiment + if update_manifest: + self._update_manifest_weight_checkpoint(self.current_exp_hash, checkpoint_file.name, step) + + # If model architecture doesn't exist in this hash directory, save a reference to where it is + self._save_architecture_reference_if_needed() + + # Persist logger queues alongside weight checkpoints + try: + self.save_logger_snapshot() + except Exception as e: + logger.debug(f"Could not save logger snapshot with checkpoint: {e}") + return checkpoint_file + except Exception as e: + logger.error(f"Failed to save model checkpoint: {e}") + return None + + def _save_architecture_reference_if_needed(self): + """Save architecture reference file if architecture doesn't exist in current hash. + + This handles the case where weights are saved to a new hash (due to HP or data changes) + but the model architecture hasn't changed. Instead of duplicating the architecture file, + we save a JSON reference pointing to the hash that contains the actual architecture. + """ + if self.current_exp_hash is None: + return + + model_dir = self.models_dir / self.current_exp_hash[8:-8] + arch_file = model_dir / f"{self.current_exp_hash[8:-8]}_architecture.pkl" + arch_ref_file = model_dir / f"{self.current_exp_hash[8:-8]}_architecture_ref.json" + + # If architecture file already exists here, no need for reference + if arch_file.exists(): + return + + # If reference file already exists, no need to create it again + if arch_ref_file.exists(): + return + + # Find the most recent hash with the same model hash that has the architecture + try: + component_hashes = self.hash_generator.get_component_hashes() + current_model_hash = component_hashes.get('model') + + if not current_model_hash: + return + + # Get all hashes with the same model hash + matching_hashes = self.get_hashes_by_component(model_hash=current_model_hash) + + # Find the most recent one that has the architecture file + for hash_candidate in sorted(matching_hashes, reverse=True): + arch_candidate = self.models_dir / hash_candidate / f"{hash_candidate}_architecture.pkl" + if arch_candidate.exists(): + # Save reference to this hash + ref_data = { + 'architecture_hash': hash_candidate, + 'current_hash': self.current_exp_hash, + 'model_hash': current_model_hash, + 'reason': 'Model architecture unchanged, reference points to hash where it is stored', + 'created': datetime.now().isoformat() + } + + os.makedirs(model_dir, exist_ok=True) + with open(arch_ref_file, 'w') as f: + json.dump(ref_data, f, indent=2) + + logger.info(f"Saved architecture reference: {self.current_exp_hash[:16]} → {hash_candidate[:16]}") + return + except Exception as e: + logger.debug(f"Could not save architecture reference: {e}") + + def save_model_architecture( + self, + model: th.nn.Module, + model_name: Optional[str] = None + ) -> Optional[Path]: + """Save full model architecture (structure + code). + + This is saved once per experiment hash when the architecture changes. + Uses dill for serialization to handle custom modules. + + Args: + model: PyTorch model + model_name: Optional name for the model + + Returns: + Path: Path to saved architecture file, or None if failed + """ + if self.current_exp_hash is None: + logger.warning("No experiment hash set. Call update_experiment_hash first.") + return None + + model_dir = self.models_dir / self.current_exp_hash[8:-8] + os.makedirs(model_dir, exist_ok=True) + arch_file = model_dir / f"{self.current_exp_hash[8:-8]}_architecture.pkl" + + # Don't overwrite if already exists + if arch_file.exists(): + logger.debug(f"Architecture already saved for {self.current_exp_hash[8:-8]}") + return arch_file + + try: + # Try dill first (better for custom classes) + if dill is not None: + with open(arch_file, 'wb') as f: + dill.dump(model, f) + else: + with open(arch_file, 'wb') as f: + pickle.dump(model, f) + + logger.info(f"Saved model architecture: {arch_file.name}") + + # Also save a text representation + arch_txt = model_dir / f"{self.current_exp_hash[8:-8]}_architecture.txt" + with open(arch_txt, 'w') as f: + f.write(str(model)) + + return arch_file + except Exception as e: + logger.error(f"Failed to save model architecture: {e}") + return None + + def save_config( + self, + config: Dict[str, Any], + config_name: str = "config" + ) -> Optional[Path]: + """Save hyperparameter configuration.""" + if self.current_exp_hash is None: + logger.warning("No experiment hash set. Call update_experiment_hash first.") + return None + + hp_hash_dir = self.hp_dir / self.current_exp_hash[:8] + os.makedirs(hp_hash_dir, exist_ok=True) + config_file = hp_hash_dir / f"{self.current_exp_hash[:8]}_{config_name}.yaml" + + try: + config_with_meta = { + 'hyperparameters': config, + 'exp_hash': self.current_exp_hash[:8], + 'last_updated': datetime.now().isoformat() + } + + with open(config_file, 'w') as f: + yaml.dump(config_with_meta, f, default_flow_style=False) + + logger.info(f"Saved config: {config_file.name}") + return config_file + except Exception as e: + logger.error(f"Failed to save config: {e}") + return None + + def save_data_snapshot(self) -> Optional[Path]: + """Save lightweight JSON snapshot of data state (sample_id, tags, deny_listed) + RNG state. + + H5 files (data.h5 and arrays.h5) are saved in parent directory (shared). + Only checkpoint-specific metadata is saved here as JSON, including random state + for reproducible data generation. + """ + if self.current_exp_hash is None: + logger.warning("No experiment hash set. Call update_experiment_hash first.") + return None + + try: + # Get dataframe manager + dfm = ledgers.get_dataframe('sample_stats') + if dfm is None: + return None + + # Trigger H5 flush to parent directory (shared) + dfm.flush_if_needed_nonblocking(force=True) + + # Extract only sample_id, tags, deny_listed for this checkpoint + df = dfm.get_df_view() + if df.empty: + return None + + if 'sample_id' not in df.columns: + df = df.reset_index() + + # Keep only checkpoint-specific columns + snapshot_cols = [ + SampleStatsEx.SAMPLE_ID.value, + SampleStatsEx.TAGS.value, + SampleStatsEx.DENY_LISTED.value + ] + available_cols = [col for col in snapshot_cols if col in df.columns] + + snapshot_df = df[available_cols] + + # Capture current RNG states for reproducibility using tool function + rng_state = capture_rng_state() + + # Convert to JSON-serializable format + snapshot_data = { + 'exp_hash': self.current_exp_hash, + 'timestamp': datetime.now().isoformat(), + 'data': snapshot_df.to_dict(orient='records'), + 'rng_state': rng_state + } + + # Capture dataloader iteration state(s) for reproducible resume (support multiple loaders) + try: + loader_states = {} + for loader_name in get_dataloaders(): + dataloader = get_dataloader(loader_name) + if dataloader is not None and hasattr(dataloader, 'capture_iteration_state'): + loader_states[loader_name] = dataloader.capture_iteration_state() + if loader_states: + snapshot_data['dataloader_iteration_state'] = loader_states + logger.debug(f"Captured dataloader iteration states: {loader_states}") + except Exception as e: + logger.debug(f"Could not capture dataloader iteration state: {e}") + + # Save to hash-specific directory + data_hash_dir = self.data_checkpoint_dir / self.current_exp_hash[-8:] + os.makedirs(data_hash_dir, exist_ok=True) + json_file = data_hash_dir / f"{self.current_exp_hash[-8:]}_data_snapshot.json" + + with open(json_file, 'w') as f: + json.dump(snapshot_data, f, indent=2, default=str) + + logger.info(f"Saved data snapshot: {json_file.name} ({len(snapshot_df)} rows) with RNG state") + return json_file + + except Exception as e: + logger.error(f"Failed to save data snapshot: {e}") + return None + + # ================= + # LOADING FUNCTIONS + # ================= + def load_latest_checkpoint( + self, + model: Optional[th.nn.Module] = None, + model_name: Optional[str] = None, + load_optimizer: bool = True, + optimizer_name: Optional[str] = None, + exp_hash: Optional[str] = None + ) -> Optional[Dict[str, Any]]: + """Load the latest model checkpoint. + + Args: + model: PyTorch model to load weights into + model_name: Name to get model from ledger + load_optimizer: Whether to load optimizer state + optimizer_name: Name to get optimizer from ledger + exp_hash: Specific experiment hash (uses current if None) + + Returns: + dict: Checkpoint data including step, metadata, etc., or None if failed + """ + target_hash = exp_hash or self.current_exp_hash + + if target_hash is None: + logger.warning("No experiment hash specified") + return None + + # Find latest checkpoint + model_dir = self.models_dir / target_hash + + if not model_dir.exists(): + logger.warning(f"No checkpoints found for hash {target_hash}") + return None + + checkpoint_files = sorted(model_dir.glob(f"{target_hash}_step_*.pt")) + + if not checkpoint_files: + logger.warning(f"No checkpoint files found in {model_dir}") + return None + + latest_checkpoint = checkpoint_files[-1] + logger.info(f"Loading checkpoint: {latest_checkpoint.name}") + + try: + checkpoint = th.load(latest_checkpoint, weights_only=False) + + # Load model state + if model is None: + try: + model = get_model(model_name) + if hasattr(model, 'get') and callable(model.get): + model = model.get() + except Exception as e: + logger.error(f"Could not get model from ledger: {e}") + return checkpoint + + if model is not None and 'model_state_dict' in checkpoint: + try: + model.load_state_dict(checkpoint['model_state_dict']) + logger.info("Loaded model state") + except Exception as e: + logger.error(f"Failed to load model state: {e}") + + # Load optimizer state + if load_optimizer and 'optimizer_state_dict' in checkpoint: + try: + optimizer = get_optimizer(optimizer_name) + if hasattr(optimizer, 'get') and callable(optimizer.get): + optimizer = optimizer.get() + if optimizer is not None: + optimizer.load_state_dict(checkpoint['optimizer_state_dict']) + logger.info("Loaded optimizer state") + except Exception as e: + logger.warning(f"Could not load optimizer state: {e}") + + return checkpoint + + except Exception as e: + logger.error(f"Failed to load checkpoint: {e}") + return None + + def list_experiment_hashes(self) -> List[str]: + """List all experiment hashes with checkpoints. + + Returns: + list: List of experiment hash strings + """ + if not self.checkpoints_dir.exists(): + return [] + + hashes = [d.name for d in self.checkpoints_dir.iterdir() if d.is_dir()] + return sorted(hashes) + + def get_checkpoint_info(self, exp_hash: Optional[str] = None) -> Dict[str, Any]: + """Get information about checkpoints for an experiment. + + Args: + exp_hash: Specific experiment hash (uses current if None) + + Returns: + dict: Information about checkpoints, configs, data backups + """ + target_hash = exp_hash or self.current_exp_hash + + if target_hash is None: + return {} + + exp_dir = self.checkpoints_dir / target_hash + + if not exp_dir.exists(): + return {} + + info = { + 'exp_hash': target_hash, + 'model_checkpoints': [], + 'architecture_saved': False, + 'configs': [], + 'data_backups': [] + } + + model_dir = exp_dir / "model" + if model_dir.exists(): + info['model_checkpoints'] = [ + f.name for f in sorted(model_dir.glob(f"{target_hash}_step_*.pt")) + ] + + # Configs + hp_dir = exp_dir / "hp" + if hp_dir.exists(): + info['configs'] = [f.name for f in hp_dir.glob("*.yaml")] + + # Data backups + data_dir = exp_dir / "data" + if data_dir.exists(): + info['data_backups'] = [f.name for f in data_dir.glob("*.h5")] + + return info + + def _save_manager_state(self): + """Save manager state (current hash, step counter, etc.)""" + state_file = self.root_log_dir / ".checkpoint_manager_state.json" + + state = { + 'current_exp_hash': self.current_exp_hash, + 'previous_exp_hash': self.previous_exp_hash, + 'step_counter': self._step_counter, + 'last_updated': datetime.now().isoformat(), + 'component_hashes': self.hash_generator.get_component_hashes(), + } + + try: + with open(state_file, 'w') as f: + json.dump(state, f, indent=2) + except Exception as e: + logger.warning(f"Failed to save manager state: {e}") + + def _load_manager_state(self): + """Load manager state if available""" + state_file = self.root_log_dir / ".checkpoint_manager_state.json" + + if not state_file.exists(): + # handled above + # No explicit state file; try to derive from manifest + manifest = self._load_manifest() + latest = manifest.get('latest_hash') + if latest: + self.current_exp_hash = latest + exp_info = manifest.get('experiments', {}).get(latest, {}) + component_hashes = { + 'hp': exp_info.get('hp_hash'), + 'model': exp_info.get('model_hash'), + 'data': exp_info.get('data_hash'), + 'combined': latest, + } + self.hash_generator.restore_hashes(component_hashes, combined_hash=latest) + logger.info(f"Derived manager state from manifest: hash={latest}") + return + + try: + with open(state_file, 'r') as f: + state = json.load(f) + + self.current_exp_hash = state.get('current_exp_hash') + self.previous_exp_hash = state.get('previous_exp_hash') + self._step_counter = state.get('step_counter', 0) + + component_hashes = state.get('component_hashes') + + # Fallback: derive component hashes from manifest when missing + if not component_hashes and self.current_exp_hash: + manifest = self._load_manifest() + exp_info = manifest.get('experiments', {}).get(self.current_exp_hash, {}) + if exp_info: + component_hashes = { + 'hp': exp_info.get('hp_hash'), + 'model': exp_info.get('model_hash'), + 'data': exp_info.get('data_hash'), + 'combined': self.current_exp_hash + } + + self.hash_generator.restore_hashes( + component_hashes, + combined_hash=self.current_exp_hash + ) + + logger.info(f"Loaded manager state: hash={self.current_exp_hash}, step={self._step_counter}") + except Exception as e: + logger.warning(f"Failed to load manager state: {e}") + + def _bootstrap_latest_state(self): + """If a current hash is known (or manifest has one), load and apply it. + + This enables auto-resume when instantiating the manager on an existing + root_log_dir without requiring an explicit load_state call by the user. + """ + target = self.current_exp_hash or self.get_latest_hash() + if not target: + return + try: + self.load_state(target) + except Exception as e: + logger.warning(f"Auto-resume failed for {target}: {e}") + + def load_checkpoint(self, + exp_hash: str, + load_model: bool = True, + load_weights: bool = True, + load_config: bool = True, + load_data: bool = True, + load_last_weights: bool = False, + force: bool = False + ) -> Dict[str, Any]: + """Load a complete checkpoint state by experiment hash. + + This method intelligently loads only the components that differ from + the current state by comparing component hashes. + + Args: + exp_hash: The 24-byte experiment hash to load (HP_MODEL_DATA) + load_model: Whether to load model architecture if different + load_weights: Whether to load model weights + load_config: Whether to load hyperparameters if different + load_data: Whether to load data state if different + load_last_weights: If True, always load the latest weights regardless of model change + force: If True, force reload of all components regardless of hash comparison + + Returns: + dict: Dictionary with keys: + - 'model': Loaded model (if changed and load_model=True) + - 'weights': Checkpoint dict with weights and metadata + - 'config': Loaded config (if changed and load_config=True) + - 'data_state': Loaded data state (if changed and load_data=True) + - 'loaded_components': Set of components that were loaded + - 'exp_hash': The experiment hash that was loaded + """ + result = { + 'model': None, + 'weights': None, + 'config': None, + 'data_state': None, + 'rng_state': None, + 'loaded_components': set(), + 'exp_hash': exp_hash + } + + # Load manifest to get component hashes + manifest = self._load_manifest() + if exp_hash not in manifest.get('experiments', {}): + logger.error(f"Experiment hash {exp_hash} not found in manifest") + return result + exp_info = manifest['experiments'][exp_hash] + target_hp_hash = exp_info.get('hp_hash') + target_model_hash = exp_info.get('model_hash') + target_data_hash = exp_info.get('data_hash') + + # Get current component hashes + current_hashes = self.hash_generator.get_component_hashes() + current_hp_hash = current_hashes.get('hp', '') + current_model_hash = current_hashes.get('model', '') + current_data_hash = current_hashes.get('data', '') + + # Logger + logger.info(f"Loading checkpoint {exp_hash[:16]}...") + logger.info(f" Target: HP={target_hp_hash} MODEL={target_model_hash} DATA={target_data_hash}") + logger.info(f" Current: HP={current_hp_hash} MODEL={current_model_hash} DATA={current_data_hash}") + + + # Load model architecture if different, or load only RNG state for reproducibility if model hash is unchanged + model_rng_loaded = False + if load_model and (target_model_hash != current_model_hash or force): + model_dir = self.models_dir / exp_hash[8:-8] + arch_ref_file = model_dir / f"{exp_hash[8:-8]}_architecture_ref.json" + + # First check if this is a reference to architecture in another hash + actual_arch_hash = exp_hash[8:-8] + if arch_ref_file.exists(): + try: + with open(arch_ref_file, 'r') as f: + ref_data = json.load(f) + actual_arch_hash = ref_data.get('architecture_hash', exp_hash[8:-8]) + logger.debug(f" Architecture reference found: pointing to hash {actual_arch_hash}") + except Exception as e: + logger.warning(f"Failed to load architecture reference: {e}") + + # Now load from actual location + actual_arch_file = self.models_dir / actual_arch_hash / f"{actual_arch_hash}_architecture.pkl" + + if actual_arch_file.exists(): + try: + with open(actual_arch_file, 'rb') as f: + result['model'] = dill.load(f) + result['loaded_components'].add('model') + logger.info(f" [OK] Loaded model architecture from hash {actual_arch_hash[:16]}") + except Exception as e: + logger.error(f" [ERROR] Failed to load model architecture: {e}") + else: + logger.warning(f" [WARNING] Model architecture file not found: {actual_arch_file}") + elif load_model and (target_model_hash == current_model_hash and not force): + # Try to load only the RNG state from the latest model checkpoint for reproducibility + model_dir = self.models_dir / exp_hash[8:-8] + checkpoint_files = sorted(model_dir.glob(f"{exp_hash}_step_*.pt")) + if not checkpoint_files: + checkpoint_files = sorted(model_dir.glob(f"{exp_hash[8:-8]}_step_*.pt")) + if checkpoint_files: + latest_checkpoint = checkpoint_files[-1] + try: + checkpoint = th.load(latest_checkpoint, weights_only=False) + rng_state = checkpoint.get('rng_state') + if rng_state: + result['rng_state'] = rng_state + model_rng_loaded = True + logger.info(f" [OK] Loaded model RNG state for reproducibility (model unchanged)") + except Exception as e: + logger.debug(f" [WARNING] Could not load model RNG state: {e}") + if not model_rng_loaded: + logger.info(f" [-] Model architecture unchanged, using current model") + else: + logger.info(f" [-] Model architecture unchanged, using current model") + + # Load model weights (always if requested) + if load_weights: + model_dir = self.models_dir / exp_hash[8:-8] + + # First, try to get the weight checkpoint from manifest for this specific experiment + checkpoint_file_to_load = None + exp_info = manifest['experiments'][exp_hash] + manifest_weight_checkpoint = exp_info.get('latest_weight_checkpoint') + + if manifest_weight_checkpoint: + checkpoint_path = model_dir / manifest_weight_checkpoint + if checkpoint_path.exists(): + checkpoint_file_to_load = checkpoint_path + logger.debug(f" Using weight checkpoint from manifest: {manifest_weight_checkpoint}") + + # Fallback: scan for weight files (old behavior for backward compatibility) + if checkpoint_file_to_load is None: + # Try new naming format first (with full exp_hash) + weight_files = sorted(model_dir.glob(f"{exp_hash}_step_*.pt")) + # Fallback to old naming format + if not weight_files: + weight_files = sorted(model_dir.glob(f"{exp_hash[8:-8]}_step_*.pt")) + + if weight_files: + checkpoint_file_to_load = weight_files[-1] # Get most recent + logger.debug(f" Using weight checkpoint from directory scan: {checkpoint_file_to_load.name}") + + if checkpoint_file_to_load: + try: + result['weights'] = th.load(checkpoint_file_to_load, weights_only=False) + result['loaded_components'].add('weights') + step = result['weights'].get('step', -1) + + # Extract RNG state from model checkpoint if available + checkpoint_rng_state = result['weights'].get('rng_state') + if checkpoint_rng_state: + result['rng_state'] = checkpoint_rng_state + logger.info(f" [OK] Loaded weights from step {step} with RNG state") + else: + logger.info(f" [OK] Loaded weights from step {step}") + + # Extract dataloader iteration state if available + dataloader_iter_state = result['weights'].get('dataloader_iteration_state') + if dataloader_iter_state: + # Normalize to mapping of loader_name -> state for backward compatibility + if isinstance(dataloader_iter_state, dict) and 'samples_yielded' in dataloader_iter_state: + iter_state_map = {'default': dataloader_iter_state} + elif isinstance(dataloader_iter_state, dict): + iter_state_map = dataloader_iter_state + else: + iter_state_map = {'default': dataloader_iter_state} + + result['dataloader_iteration_state'] = iter_state_map + logger.debug(f" [OK] Found dataloader iteration state(s): {iter_state_map}") + except Exception as e: + logger.error(f" [ERROR] Failed to load weights: {e}") + else: + logger.warning(f" [WARNING] No weight files found for {exp_hash[8:-8]}") + + # Load config if different + if load_config and (target_hp_hash != current_hp_hash or force): + hp_dir = self.hp_dir / exp_hash[:8] + config_file = hp_dir / f"{exp_hash[:8]}_config.yaml" + + if config_file.exists(): + try: + with open(config_file, 'r') as f: + config_data = yaml.safe_load(f) + result['config'] = config_data.get('hyperparameters', config_data) + result['loaded_components'].add('config') + logger.info(f" [OK] Loaded config (hash changed)") + except Exception as e: + logger.error(f" [ERROR] Failed to load config: {e}") + else: + logger.warning(f" [WARNING] Config file not found: {config_file}") + else: + logger.info(f" [-] Config unchanged, using current config") + + # Load data snapshot if different, or if only RNG state changed (for reproducibility) + if load_data: + data_dir = self.data_checkpoint_dir / exp_hash[-8:] + json_file = data_dir / f"{exp_hash[-8:]}_data_snapshot.json" + + # Always try to load RNG state for reproducibility, even if data hash is unchanged + load_data_snapshot = (target_data_hash != current_data_hash or force) + load_rng_only = (target_data_hash == current_data_hash and not force) + + if json_file.exists(): + try: + with open(json_file, 'r') as f: + snapshot_data = json.load(f) + + rng_state = snapshot_data.get('rng_state', {}) + + if load_data_snapshot: + snapshot_df = pd.DataFrame(snapshot_data.get('data', [])) + if not snapshot_df.empty: + result['data_state'] = {'snapshot': snapshot_df} + result['loaded_components'].add('data') + if rng_state: + result['rng_state'] = rng_state + logger.info(f" [OK] Loaded data snapshot ({len(snapshot_df)} rows) with RNG state") + else: + logger.info(f" [OK] Loaded data snapshot ({len(snapshot_df)} rows)") + elif load_rng_only and rng_state: + # Only RNG state is needed for reproducibility + result['rng_state'] = rng_state + logger.info(f" [OK] Loaded RNG state for reproducibility (data unchanged)") + else: + logger.info(f" [-] Data state unchanged, using current data") + except Exception as e: + logger.error(f" [ERROR] Failed to load data snapshot: {e}") + else: + logger.warning(f" [WARNING] Data snapshot file not found: {json_file}") + + logger.info(f"Loaded components: {result['loaded_components']}") + return result + + def load_state(self, exp_hash: str, load_last_weights: bool = False) -> bool: + """Load and apply a complete checkpoint state by experiment hash. + + This method loads all components and updates the system state in-place: + - Updates model in ledger (architecture + weights) + - Updates config in ledger + - Updates dataframe manager with loaded data + - Updates current experiment hash + + Args: + exp_hash: The 24-byte experiment hash to load and apply + load_last_weights: If True, always load the latest weights regardless of model change + + Returns: + bool: True if state was successfully loaded and applied + """ + logger.info(f"\n{'='*60}") + logger.info(f"Loading and applying state: {exp_hash[:16]}...") + logger.info(f"{'='*60}") + + # Load checkpoint data + checkpoint_data = self.load_checkpoint( + exp_hash=exp_hash, + load_model=True, + load_weights=True, + load_config=True, + load_data=True, + load_last_weights=load_last_weights + ) + + if not checkpoint_data['loaded_components']: + logger.warning("No components were loaded") + return False + + success = True + + # Apply model (architecture + weights) + if 'model' in checkpoint_data['loaded_components']: + try: + model = checkpoint_data['model'] + + # Register in ledger + ledgers.register_model(ledgers.resolve_hp_name(), model) + + # Set Model Training Guard + guard_training_context.model = model # Train + guard_testing_context.model = model # Eval + except Exception as e: + logger.error(f"[ERROR] Failed to apply model: {e}") + success = False + elif 'weights' in checkpoint_data['loaded_components']: + # Only weights changed, apply to existing model + try: + model = ledgers.get_model() + weights = checkpoint_data['weights'] + if model and weights and 'model_state_dict' in weights: + model.load_state_dict(weights['model_state_dict']) + step = weights.get('step', -1) + logger.info(f"[OK] Applied weights to existing model (step {step})") + + # Set Model Training Guard + guard_training_context.model = model # Train + guard_testing_context.model = model # Eval + except Exception as e: + logger.error(f"[ERROR] Failed to apply weights: {e}") + success = False + + # Apply config + if 'config' in checkpoint_data['loaded_components']: + try: + config = checkpoint_data['config'] + ledgers.register_hyperparams(ledgers.resolve_hp_name(), config) + logger.info(f"[OK] Applied hyperparameters config") + except Exception as e: + logger.error(f"[ERROR] Failed to apply config: {e}") + success = False + + # Apply data (merge snapshot columns into current dataframe) + if 'data' in checkpoint_data['loaded_components']: + try: + data_state = checkpoint_data.get('data_state', {}) + snapshot_df = data_state.get('snapshot') + + if snapshot_df is not None and not snapshot_df.empty: + dfm = ledgers.get_dataframe('sample_stats') + if dfm is not None: + # Set index if needed + if 'sample_id' in snapshot_df.columns: + snapshot_df = snapshot_df.set_index('sample_id') + + # Merge only the checkpoint-specific columns (tags, deny_listed) + # This updates existing rows without replacing all data + dfm.upsert_df(snapshot_df, force_flush=True) + logger.info(f"[OK] Applied data snapshot ({len(snapshot_df)} rows)") + except Exception as e: + logger.error(f"[ERROR] Failed to apply data: {e}") + success = False + + # Restore RNG state if provided and not already restored + if checkpoint_data.get('rng_state'): + try: + restore_rng_state(checkpoint_data['rng_state']) + logger.debug(f"Restored RNG state from checkpoint") + + # Reset dataloaders iterators to ensure reproducibility + for loader_name in ledgers.get_dataloaders(): + loader = ledgers.get_dataloader(loader_name) + + if loader is not None: + # Resume loader state + if hasattr(loader, 'reset_iterator') and callable(loader.reset_iterator): + loader.reset_iterator() + logger.debug(f"Reset iterator for dataloader: {loader}") + + # Restore RNG state again after resetting dataloaders + restore_rng_state(checkpoint_data['rng_state']) + logger.debug(f"Restored RNG state from checkpoint") + + except Exception as e: + logger.error(f"[ERROR] Failed to restore RNG state: {e}") + pause_ctrl.pause() + success = False + + # Restore dataloader iteration state if provided + if checkpoint_data.get('dataloader_iteration_state'): + try: + iter_state_raw = checkpoint_data['dataloader_iteration_state'] + + # Normalize to mapping loader_name -> state for backward compatibility + if isinstance(iter_state_raw, dict) and 'samples_yielded' in iter_state_raw: + state_map = {'default': iter_state_raw} + elif isinstance(iter_state_raw, dict): + state_map = iter_state_raw + else: + state_map = {'default': iter_state_raw} + + restored_any = False + for loader_name in ledgers.get_dataloaders(): + loader = ledgers.get_dataloader(loader_name) + if loader is None or not hasattr(loader, 'restore_iteration_state'): + continue + + state_for_loader = state_map.get(loader_name) or state_map.get('default') + if state_for_loader: + try: + loader.restore_iteration_state(state_for_loader) + # Resume loader state + if hasattr(loader, 'reset_iterator') and callable(loader.reset_iterator): + loader.reset_iterator() + logger.debug(f"Reset iterator for dataloader: {loader}") + logger.info(f"[OK] Restored dataloader iteration state for {loader_name}: {state_for_loader}") + restored_any = True + except Exception as inner_e: + logger.warning(f"[WARNING] Failed to restore iteration state for {loader_name}: {inner_e}") + + if not restored_any: + logger.warning("No dataloader iteration state could be applied to registered loaders") + except Exception as e: + logger.error(f"[ERROR] Failed to restore dataloader iteration state: {e}") + success = False + + # Restore logger snapshot for this experiment if available + try: + self.load_logger_snapshot(exp_hash) + except Exception as e: + logger.warning(f"Failed to restore logger snapshot for {exp_hash}: {e}") + + # Update current experiment hash + if success: + old_hash = self.current_exp_hash + self.current_exp_hash = exp_hash + self.previous_exp_hash = old_hash + + # Keep hash generator in sync with loaded experiment + manifest = self._load_manifest() + exp_info = manifest.get('experiments', {}).get(exp_hash, {}) + component_hashes = { + 'hp': exp_info.get('hp_hash'), + 'model': exp_info.get('model_hash'), + 'data': exp_info.get('data_hash'), + 'combined': exp_hash + } + self.hash_generator.restore_hashes(component_hashes, combined_hash=exp_hash) + + self._save_manager_state() + logger.info(f"\n[OK] Successfully loaded and applied state: {exp_hash[:16]}") + else: + logger.warning(f"\n[WARNING] State loaded with errors") + + logger.info(f"{'='*60}\n") + return success diff --git a/weightslab/components/experiment_hash.py b/weightslab/components/experiment_hash.py new file mode 100644 index 00000000..5b4f8c0c --- /dev/null +++ b/weightslab/components/experiment_hash.py @@ -0,0 +1,328 @@ +""" +Experiment Hash Generation Module + +This module generates stable, deterministic hashes for experiment tracking. +The hash is based on three key components: + 1. Model architecture (structure and layer configuration) - 8 bytes + 2. Hyperparameters (config values) - 8 bytes + 3. Data state (UIDs, discard status, tags) - 8 bytes + +Combined into a 24-byte hash that allows tracking what changed between versions. +""" + +import hashlib +import json +import logging +from typing import Any, Dict, List, Optional, Set +from pathlib import Path + +import torch as th + + +logger = logging.getLogger(__name__) + + +class ExperimentHashGenerator: + """Generates deterministic hashes for experiment tracking. + + Computes three separate 8-byte hashes: + - Hyperparameters hash (learning rate, batch size, etc.) + - Model architecture hash (layers, parameters, structure) + - Data hash (UIDs, discard status, tags) + + These are combined into a final 24-byte hash in order: HP_MODEL_DATA + This allows tracking what changed between experiment versions. + + Attributes: + _last_hash (str): The most recently generated combined hash (24 chars) + _last_hp_hash (str): Hash of the hyperparameters (8 chars) + _last_model_hash (str): Hash of the model architecture (8 chars) + _last_data_hash (str): Hash of the data state (8 chars) + """ + + def __init__(self): + self._last_hash: Optional[str] = None + self._last_hp_hash: Optional[str] = None + self._last_model_hash: Optional[str] = None + self._last_data_hash: Optional[str] = None + + def generate_hash( + self, + model: Optional[th.nn.Module] = None, + config: Optional[Dict[str, Any]] = None, + data_state: Optional[Dict[str, Any]] = None + ) -> str: + """Generate a unique hash for the current experiment configuration. + + Computes three separate 8-byte hashes and combines them into a 24-byte hash. + This allows tracking which component changed (model, config, or data). + + Args: + model: PyTorch model to hash (architecture only, not weights) + config: Dictionary of hyperparameters + data_state: Dictionary with 'uids', 'discarded', 'tags' for data samples + + Returns: + str: A 24-character hexadecimal hash string (8 + 8 + 8) + """ + # Generate individual 8-byte hashes + hp_hash = self._hash_config(config) if config is not None else "00000000" + model_hash = self._hash_model(model) if model is not None else "00000000" + data_hash = self._hash_data_state(data_state) if data_state is not None else "00000000" + + # Combine into 24-byte hash: HP (8) + MODEL (8) + DATA (8) + final_hash = f"{hp_hash}{model_hash}{data_hash}" + + # Store for comparison + self._last_hash = final_hash + self._last_hp_hash = hp_hash + self._last_model_hash = model_hash + self._last_data_hash = data_hash + + logger.info(f"Generated experiment hash: {final_hash}") + logger.debug(f" HP hash: {hp_hash}") + logger.debug(f" Model hash: {model_hash}") + logger.debug(f" Data hash: {data_hash}") + + return final_hash + + def has_changed( + self, + model: Optional[th.nn.Module] = None, + config: Optional[Dict[str, Any]] = None, + data_state: Optional[Dict[str, Any]] = None, + force: bool = False + ) -> tuple[bool, Set[str]]: + """Check if the experiment configuration has changed. + + Checks all three components: model architecture, hyperparameters, and data. + + Args: + model: PyTorch model to check + config: Dictionary of hyperparameters to check + data_state: Dictionary with data state (uids, discarded, tags) + + Returns: + tuple: (has_changed: bool, changed_components: Set[str]) + where changed_components can contain 'model', 'config', 'data' + """ + changed_components = set() + + # Check HP + if config is not None: + hp_hash = self._hash_config(config) + if hp_hash != self._last_hp_hash or force: + changed_components.add('hp') + + # Check model + if model is not None: + model_hash = self._hash_model(model) + if model_hash != self._last_model_hash or force: + changed_components.add('model') + + # Check data + if data_state is not None: + data_hash = self._hash_data_state(data_state) + if data_hash != self._last_data_hash or force: + changed_components.add('data') + + has_changed = len(changed_components) > 0 + + if has_changed: + logger.info(f"Experiment configuration changed: {changed_components}") + + return has_changed, changed_components + + def _hash_model(self, model: th.nn.Module) -> str: + """Generate a hash from model architecture. + + This captures the model structure (layer types, parameters, connections) + but not the actual weights, so the same architecture always produces + the same hash. + + Args: + model: PyTorch model + + Returns: + str: Hash of model architecture (8 bytes) + + TODO (GP): Hash should be generated directly from the model class and computed on demand. Same for data and HP. + Maybe later neurons tracking and values in the hash. + """ + try: + # Get model architecture info + arch_info = [] + + # Model class name + arch_info.append(f"class:{model.__class__.__name__}") + + # Layer structure + for name, module in model.named_modules(): + # Remove these trackers from hash + if 'train_dataset_tracker' in name or 'eval_dataset_tracker' in name: + continue + if name: # Skip root module + module_info = f"{name}:{module.__class__.__name__}" + + # Add key parameters for common layer types + if isinstance(module, th.nn.Module) and hasattr(module, 'in_neurons') and hasattr(module, 'out_neurons'): + module_info += f"_in{module.in_neurons}_out{module.out_neurons}" + if hasattr(module, 'operation_age'): + for op_type, age in module.operation_age.items(): + module_info += f"_{op_type}->{age}" + + arch_info.append(module_info) + + # Create hash from architecture description (8 bytes = 8 chars hex) + arch_str = "|".join(sorted(arch_info)) + return hashlib.sha256(arch_str.encode()).hexdigest()[:8] + + except Exception as e: + logger.warning(f"Failed to hash model architecture: {e}") + # Fallback: use model repr (8 bytes) + try: + return hashlib.sha256(str(model).encode()).hexdigest()[:8] + except Exception: + return "00000000" + + def _hash_config(self, config: Dict[str, Any]) -> str: + """Generate a hash from hyperparameters configuration. + + Args: + config: Dictionary of hyperparameters + + Returns: + str: Hash of configuration (8 bytes) + """ + # Remove random state from config, i.e., root log dir as can be generated randomly + config_cp = config.copy() + config_cp.pop('root_log_dir', None) + config_cp.pop('is_training', None) + try: + # Sort keys for deterministic hashing + # Convert to JSON string for stable representation + config_str = json.dumps(config_cp, sort_keys=True, default=str) + return hashlib.sha256(config_str.encode()).hexdigest()[:8] + except Exception as e: + logger.warning(f"Failed to hash config: {e}") + return "00000000" + + def _hash_data_state(self, data_state: Dict[str, Any]) -> str: + """Generate a hash from data state (UIDs, discard status, tags). + + Args: + data_state: Dictionary with 'uids', 'discarded', 'tags' + - uids: List of sample UIDs + - discarded: Set of discarded UIDs + - tags: Dict mapping UID to list of tags + + Returns: + str: Hash of data state (8 bytes) + """ + try: + # Extract components + uids = list(data_state.get('discarded', dict()).keys()) + discarded = data_state.get('discarded', dict()) + tags = data_state.get('tags', {}) + + # Create deterministic representation + # Sort UIDs and include discard status and tags + data_info = [] + for uid in sorted(uids): + is_discarded = discarded[uid] + uid_tags = sorted(tags.get(uid, [])) + data_info.append(f"{uid}:d{int(is_discarded)}:t{','.join(uid_tags)}") + + data_str = "|".join(data_info) + return hashlib.sha256(data_str.encode()).hexdigest()[:8] + except Exception as e: + logger.warning(f"Failed to hash data state: {e}") + return "00000000" + + def get_last_hash(self) -> Optional[str]: + """Get the most recently generated hash. + + Returns: + str or None: Last generated hash (24 bytes), or None if no hash generated yet + """ + return self._last_hash + + def get_component_hashes(self) -> Dict[str, Optional[str]]: + """Get individual component hashes. + + Returns: + dict: Dictionary with 'hp', 'model', 'data' hash values (8 bytes each) + and 'combined' (24 bytes total) + """ + return { + 'hp': self._last_hp_hash, + 'model': self._last_model_hash, + 'data': self._last_data_hash, + 'combined': self._last_hash + } + + def restore_hashes( + self, + component_hashes: Optional[Dict[str, Optional[str]]] = None, + combined_hash: Optional[str] = None + ) -> None: + """Restore last known hashes so change detection stays consistent. + + This is used when reloading manager state to keep the hash generator + in sync with the previously computed hashes. If a combined hash is + provided, it is split into component hashes when they are missing. + """ + hashes = component_hashes or {} + + hp_hash = hashes.get('hp') if isinstance(hashes, dict) else None + model_hash = hashes.get('model') if isinstance(hashes, dict) else None + data_hash = hashes.get('data') if isinstance(hashes, dict) else None + + combined = (hashes.get('combined') if isinstance(hashes, dict) else None) or combined_hash + + # If we have a combined hash, use it to fill missing components + if combined and len(str(combined)) >= 24: + combined_str = str(combined)[:24] + hp_hash = hp_hash or combined_str[0:8] + model_hash = model_hash or combined_str[8:16] + data_hash = data_hash or combined_str[16:24] + combined = combined_str + + # If components are present but combined is missing, rebuild combined + if not combined and hp_hash and model_hash and data_hash: + combined = f"{hp_hash}{model_hash}{data_hash}" + + self._last_hp_hash = hp_hash + self._last_model_hash = model_hash + self._last_data_hash = data_hash + self._last_hash = combined + + logger.debug( + f"Restored hashes hp={hp_hash}, model={model_hash}, data={data_hash}, combined={combined}" + ) + + def compare_hashes(self, hash1: str, hash2: str) -> Set[str]: + """Compare two 24-byte hashes and identify what changed. + + Args: + hash1: First hash (24 chars) + hash2: Second hash (24 chars) + + Returns: + Set of changed components: 'model', 'config', 'data' + """ + if len(hash1) != 24 or len(hash2) != 24: + logger.warning(f"Invalid hash lengths: {len(hash1)}, {len(hash2)}") + return set() + + changed = set() + + # Compare each 8-byte segment (HP_MODEL_DATA) + if hash1[0:8] != hash2[0:8]: + changed.add('hp') + if hash1[8:16] != hash2[8:16]: + changed.add('model') + if hash1[16:24] != hash2[16:24]: + changed.add('data') + + return changed diff --git a/weightslab/components/global_monitoring.py b/weightslab/components/global_monitoring.py index 1e297245..b6f4cc3b 100644 --- a/weightslab/components/global_monitoring.py +++ b/weightslab/components/global_monitoring.py @@ -5,7 +5,7 @@ import time import logging -from weightslab.backend.ledgers import get_hyperparams, set_hyperparam, resolve_hp_name +from weightslab.backend.ledgers import get_hyperparams, set_hyperparam, resolve_hp_name, get_checkpoint_manager, get_optimizers, get_optimizer from weightslab.components.tracking import TrackingMode @@ -23,23 +23,49 @@ class PauseController: def __init__(self): self._event = Event() self._event.clear() - # self._event.set() # Set by default so training starts running + + # Get checkpoint manager instance + self.checkpoint_manager = None def wait_if_paused(self): # Called from main thread / model forward. Blocks if paused. self._event.wait() # releases GIL while waiting def pause(self): - logger.info('\nTraining paused.') self._event.clear() + logger.info('\nTraining paused.') def resume(self): - logger.info('\nTraining resumed.') - self._event.set() + # On resume, first dump any pending changes to checkpoint manager + if self.checkpoint_manager is None: + self.checkpoint_manager = get_checkpoint_manager() + if self.checkpoint_manager is not None: + self.checkpoint_manager.update_experiment_hash(firsttime=True) + self.checkpoint_manager.dump_pending_changes() + + # Then resume execution + if self._is_hash_computed(): + self._event.set() + logger.info(f'\nTraining resumed as modules hashes have been computed: {self.checkpoint_manager.hash_by_module}.') + return True + else: + logger.warning(f'Cannot resume training: experiment hash not computed yet for every modules {self.checkpoint_manager.hash_by_module}.') + return False def is_paused(self): return not self._event.is_set() + def _get_checkpoint_manager(self): + if self.checkpoint_manager is None: + self.checkpoint_manager = get_checkpoint_manager() + + def _is_hash_computed(self): + self._get_checkpoint_manager() + if self.checkpoint_manager is None: + return False + fl = self.checkpoint_manager.hash_by_module[0] != "00000000" and self.checkpoint_manager.hash_by_module[1] != "00000000" and self.checkpoint_manager.hash_by_module[2] != "00000000" + + return fl # Global pause controller instance pause_controller = PauseController() @@ -113,59 +139,6 @@ def __exit__(self, exc_type: Any, exc_value: Any, traceback: Any) -> bool: self.architecture_guard.__exit__(exc_type, exc_value, traceback) - # Use provided op_context if present, otherwise fall back to module-level - ctx = op_context if getattr(self, 'op_context', None) is not None else op_context - with ctx: - # decrement training steps and store result in ledgered hyperparams - try: - # resolve a sensible hyperparam set name (reuse helper in this module) - name = resolve_hp_name() - if name is not None: - try: - hp_handle = get_hyperparams(name) - except Exception: - hp_handle = None - - try: - if hp_handle is None: - raise RuntimeError('no hyperparams') - # unwrap proxy if present - if hasattr(hp_handle, 'get') and not isinstance(hp_handle, dict): - hp = hp_handle.get() - else: - hp = hp_handle - - if not isinstance(hp, dict): - raise RuntimeError('hyperparams not a dict') - - cur = hp.get('training_steps_to_do', 0) - try: - cur_int = int(cur) - except Exception: - cur_int = 0 - new = max(0, cur_int - 1) - - # try ledger API first - try: - set_hyperparam(name, 'training_steps_to_do', new) - except Exception: - # best-effort fallback: update dict directly - try: - hp['training_steps_to_do'] = new - except Exception: - pass - except Exception: - # swallow errors - don't let monitoring break training - pass - except Exception: - pass - - # Auto-increment step count for UI progress - if self.for_training and self.model is not None: - if not hasattr(self.model, 'current_step'): - self.model.current_step = 0 - self.model.current_step += 1 - return False @@ -181,9 +154,10 @@ def __exit__(self, exc_type: Any, exc_value: Any, traceback: Any) -> bool: # - If controller is paused/resumed externally, update ledger `is_training` to match. _pause_sync_thread_started = False - +checkpoint_manager = get_checkpoint_manager() def _pause_hp_sync_loop(poll_interval: float = 0.5): + firstresume = True while True: try: name = resolve_hp_name() @@ -220,7 +194,8 @@ def _pause_hp_sync_loop(poll_interval: float = 0.5): # Drive controller from ledger when ledger explicitly sets the flag if isinstance(hp_is_training, bool): if controller_paused and hp_is_training: - pause_controller.resume() + resumed = pause_controller.resume() + firstresume = False if resumed else True elif controller_running and not hp_is_training: pause_controller.pause() @@ -228,11 +203,12 @@ def _pause_hp_sync_loop(poll_interval: float = 0.5): controller_paused = pause_controller.is_paused() # Propagate controller state back to ledger if it differs - if controller_paused: + if controller_paused and not firstresume: set_hyperparam(name, 'is_training', False) - except Exception: + except Exception as e: # swallow to keep thread alive + logger.debug(f"Exception in pause-hp sync loop: {e}") pass time.sleep(poll_interval) diff --git a/weightslab/data/data_samples_with_ops.py b/weightslab/data/data_samples_with_ops.py index addb49c5..064ba97b 100644 --- a/weightslab/data/data_samples_with_ops.py +++ b/weightslab/data/data_samples_with_ops.py @@ -12,6 +12,7 @@ from enum import Enum from typing import Callable, Any, Set, Dict, Optional from torch.utils.data import Dataset, Subset +from weightslab.backend import ledgers from weightslab.utils.tools import array_id_2bytes from weightslab.data.h5_dataframe_store import H5DataFrameStore from weightslab.trainer.services.service_utils import load_label @@ -28,6 +29,8 @@ SAMPLE_STATS_ALL, ) +from weightslab.components.checkpoint_manager_v2 import CheckpointManagerV2 + # Global logger logger = logging.getLogger(__name__) @@ -166,9 +169,13 @@ def __init__( logger.info(f"[DataSampleTrackingWrapper] Using temporary directory {self._root_log_dir} for H5 persistence. Please copy final results in a safe location after training.") if self._enable_h5_persistence and self._root_log_dir: + # Store H5 files in PARENT directory (shared across all experiment hashes) + # Only checkpoint-specific JSON files go in hash directories data_dir = self._root_log_dir / "checkpoints" / "data" data_dir.mkdir(parents=True, exist_ok=True) - self._h5_path = data_dir / "data_with_ops.h5" + + # Use shared data.h5 file (not hash-specific) + self._h5_path = data_dir / "data.h5" logger.info(f"[DataSampleTrackingWrapper] H5 persistence enabled at {self._h5_path}") # If no shared store provided, create one pointing to the same path diff --git a/weightslab/data/dataframe_manager.py b/weightslab/data/dataframe_manager.py index 80ed39ac..edba9dcc 100644 --- a/weightslab/data/dataframe_manager.py +++ b/weightslab/data/dataframe_manager.py @@ -1,3 +1,4 @@ +import os import threading import logging import traceback @@ -11,6 +12,8 @@ from weightslab.data.h5_dataframe_store import H5DataFrameStore from weightslab.data.h5_array_store import H5ArrayStore +from weightslab.backend import ledgers +from weightslab.components import CheckpointManagerV2 from weightslab.data.array_proxy import ArrayH5Proxy, convert_dataframe_to_proxies from weightslab.data.data_utils import _filter_columns_by_patterns from weightslab.backend.ledgers import get_dataloaders, get_dataloader @@ -71,8 +74,9 @@ def set_store(self, store: H5DataFrameStore): with self._lock: if self._store is None and self._enable_h5_persistence: self._store = store - # Auto-create array store in same directory + # Auto-create array store in SAME directory (shared, both in parent) if self._array_store is None: + # data.h5 is already in checkpoints/data/, so arrays.h5 goes there too array_path = store.get_path().parent / "arrays.h5" self._array_store = H5ArrayStore(array_path) @@ -141,7 +145,6 @@ def _load_existing_data(self, origin: str = None, autoload_arrays: bool | list | else: logger.warning(f"[LedgeredDataFrameManager] Loaded data missing 'sample_id' column for origin={origin}. Skipping load.") - def upsert_df(self, df_local: List | pd.DataFrame, origin: str = None, force_flush: bool = False): if df_local is None or (isinstance(df_local, pd.DataFrame) and df_local.empty) or len(df_local) == 0: return @@ -167,18 +170,16 @@ def upsert_df(self, df_local: List | pd.DataFrame, origin: str = None, force_flu with self._lock: # Align columns - all_cols = self._df.columns.union(df_norm.columns) + all_cols = df_norm.columns if self._df.empty: self._df = df_norm.reindex(columns=all_cols) return - if len(all_cols) != len(self._df.columns): - self._df = self._df.reindex(columns=all_cols) - if len(all_cols) != len(df_norm.columns): - df_norm = df_norm.reindex(columns=all_cols) # Right-preferred upsert: df_norm overrides existing, adds new rows - # Override existing rows where sample_id matches - self._df.update(df_norm) + # Only update columns present in df_norm, keep other columns/values from self._df + existing_idx = df_norm.index.intersection(self._df.index) + if len(existing_idx) > 0: + self._df.loc[existing_idx, all_cols] = df_norm.loc[existing_idx, all_cols] # Append rows that do not exist yet missing_idx = df_norm.index.difference(self._df.index) @@ -308,7 +309,7 @@ def is_meaningful(v): # Build all records BEFORE acquiring lock (faster) records_to_add = [] for i, sid in enumerate(sample_ids): - sample_id = int(sid) + sample_id = int(sid) if not isinstance(sid, np.ndarray) else int(sid[0]) # Build record incrementally - keep numpy arrays as-is for speed rec: Dict[str, Any] = { @@ -424,10 +425,8 @@ def get_df_view(self, column: str = None, limit: int = -1, copy: bool = False) - with self._lock: if self._df.empty: return pd.DataFrame() - if column is not None and column in self._df.columns: - mask = self._df[column] == column - # Return view of matching rows - subset = self._df.loc[mask] + if column is not None and ((not isinstance(column, (list, set, tuple)) and column in self._df.columns) or (isinstance(column, (list, set, tuple)))): + subset = self._df[column] else: subset = self._df if limit > 0: @@ -836,7 +835,7 @@ def flush_async(self): self._flush_queue_count += 1 self._flush_event.set() # Wake thread immediately - def flush_if_needed_nonblocking(self, force: bool = False): + def flush_if_needed_nonblocking(self, force: bool = False, force_flush_h5: bool = False): """Non-blocking flush - if can't acquire lock immediately, defer to next cycle.""" # Try to acquire buffer lock with timeout if not self._buffer_lock.acquire(blocking=False): @@ -850,6 +849,8 @@ def flush_if_needed_nonblocking(self, force: bool = False): buffered = list(self._buffer.values()) self._buffer = {} finally: + if force_flush_h5: + self._flush_to_h5_if_needed(force=force) self._buffer_lock.release() # Apply records outside buffer lock diff --git a/weightslab/data/h5_dataframe_store.py b/weightslab/data/h5_dataframe_store.py index c1b1eefb..056f4daf 100644 --- a/weightslab/data/h5_dataframe_store.py +++ b/weightslab/data/h5_dataframe_store.py @@ -99,6 +99,8 @@ class H5DataFrameStore: - Keeps a stable schema by treating `sample_id` as the index and always tagging rows with `origin`. - Provides small helpers for slice-based reads used by the DataService. + + TODO (GP): Refactor both h5 functions into common utility module first. """ def __init__(self, path: Union[str, Path], key_prefix: str = "stats", lock_timeout: float = 10.0, poll_interval: float = 0.1): diff --git a/weightslab/examples/data/ws-classification/config.yaml b/weightslab/examples/data/ws-classification/config.yaml index f5172cdd..ca894fa6 100644 --- a/weightslab/examples/data/ws-classification/config.yaml +++ b/weightslab/examples/data/ws-classification/config.yaml @@ -16,8 +16,8 @@ ledger_flush_max_rows: 750 ledger_flush_interval: 30.0 # Configure clients -serving_grpc: false -serving_cli: false +serving_grpc: true +serving_cli: true # DataLoader parameters data: diff --git a/weightslab/examples/data/ws-classification/main.py b/weightslab/examples/data/ws-classification/main.py index 8d73a0d8..229ab399 100644 --- a/weightslab/examples/data/ws-classification/main.py +++ b/weightslab/examples/data/ws-classification/main.py @@ -104,7 +104,6 @@ def test(loader, model, criterion_mlt, metric_mlt, device, test_loader_len): preds=preds, ) - loss = losses / test_loader_len metric = metric_mlt.compute() * 100 @@ -170,7 +169,7 @@ def test(loader, model, criterion_mlt, metric_mlt, device, test_loader_len): # Data (MNIST train/val/test) _full_train_dataset = datasets.MNIST( - root=os.path.join(parameters["root_log_dir"], "data"), + root=os.path.join(r'C:/Users/GuillaumePelluet/Desktop/trash/cls_usecase/', "data"), train=True, download=True, transform=transforms.Compose( @@ -180,7 +179,7 @@ def test(loader, model, criterion_mlt, metric_mlt, device, test_loader_len): ), ) _test_dataset = datasets.MNIST( - root=os.path.join(parameters["root_log_dir"], "data"), + root=os.path.join(r'C:/Users/GuillaumePelluet/Desktop/trash/cls_usecase/', "data"), train=False, download=True, transform=transforms.Compose( diff --git a/weightslab/models/model_with_ops.py b/weightslab/models/model_with_ops.py index 63b31b32..49459a70 100644 --- a/weightslab/models/model_with_ops.py +++ b/weightslab/models/model_with_ops.py @@ -124,6 +124,7 @@ def to(self, device, dtype=None, non_blocking=False, **kwargs): super().to(device, dtype, non_blocking, **kwargs) for layer in self.layers: layer.to(device, dtype, non_blocking, **kwargs) + return self def maybe_update_age(self, tracked_input: th.Tensor): if self.tracking_mode != TrackingMode.TRAIN: @@ -477,19 +478,38 @@ def _operate( ) def state_dict(self, destination: Optional[Dict[str, Any]] = None, prefix: str = '', keep_vars: bool = False) -> Dict[str, Any]: - state = super().state_dict(**{'destination': destination, 'prefix': prefix, 'keep_vars': keep_vars}) - state[prefix + 'seen_samples'] = self.seen_samples - state[prefix + 'current_step'] = self.current_step - state[prefix + 'tracking_mode'] = self.tracking_mode - return state + state_dict = super().state_dict(**{'destination': destination, 'prefix': prefix, 'keep_vars': keep_vars}) + state_dict[prefix + 'seen_samples'] = self.seen_samples + state_dict[prefix + 'current_step'] = self.current_step + state_dict[prefix + 'tracking_mode'] = self.tracking_mode + return state_dict def load_state_dict( - self, state_dict, strict, assign=True, **kwargs): - self.seen_samples = state_dict['seen_samples'] - self.current_step = state_dict.get('current_step', 0) - self.tracking_mode = state_dict['tracking_mode'] - super().load_state_dict( - state_dict, strict=strict, assign=assign, **kwargs) + self, state_dict, strict=True, assign=True, **kwargs): + self.seen_samples = state_dict.pop('seen_samples', 0) + self.current_step = state_dict.pop('current_step', 0) + self.tracking_mode = state_dict.pop('tracking_mode', 0) + + # Preprocess trackers + # TODO (GP): better way to handle this? Maybe load the trackers + state_dict = { + k: v for k, v in state_dict.items() + if not '_dataset_tracker' in k + } + try: + super().load_state_dict( + state_dict, strict=strict, assign=assign, **kwargs) + except Exception: + # If the state dict comes from a wrapper (e.g., ModelInterface) it may + # include a leading "model." prefix. Strip that prefix for matching. + remapped = {} + for k, v in state_dict.items(): + if k.startswith('model.'): + remapped[k[len('model.'):]] = v + else: + remapped[k] = v + super().load_state_dict( + remapped, strict=strict, assign=assign, **kwargs) def forward(self, tensor: th.Tensor, diff --git a/weightslab/modules/modules_with_ops.py b/weightslab/modules/modules_with_ops.py index facf7ba9..d38264c2 100644 --- a/weightslab/modules/modules_with_ops.py +++ b/weightslab/modules/modules_with_ops.py @@ -46,6 +46,7 @@ def __init__( self.module_name = module_name self.device = device self.tracking_mode = TrackingMode.DISABLED + self.operation_age = {op.name: 0 for op in ArchitectureNeuronsOpType} # keep track of all operations performed # IN/OUT neurons indexing & mapping dictionary self.src_to_dst_mapping_tnsrs = {} @@ -115,14 +116,14 @@ def get_neurons(self, attr_name: str) -> int: f"Accessing '{attr_name}' before calling " + "_initialize_neuron_attributes." ) - + val = getattr(self, attr_name) if val is None and getattr(self, 'wl_same_flag', False): # For pass-through layers, falling back to the other dimension if one is None other_attr = 'out_neurons' if attr_name == 'in_neurons' else 'in_neurons' if hasattr(self, other_attr): val = getattr(self, other_attr) - + return self.get_neurons_value(val) def set_neurons( @@ -668,6 +669,9 @@ def operate( **kwargs ) + # Set Operation flag + self.operation_age[op_type.name] += 1 + # ------------------ # Neurons Operations def _process_input_neurons_index( diff --git a/weightslab/proto/experiment_service.proto b/weightslab/proto/experiment_service.proto index db053595..1142cbd3 100644 --- a/weightslab/proto/experiment_service.proto +++ b/weightslab/proto/experiment_service.proto @@ -2,7 +2,7 @@ syntax = "proto3"; service ExperimentService { - rpc StreamStatus (Empty) returns (stream TrainingStatusEx); + rpc GetLatestLoggerData (GetLatestLoggerDataRequest) returns (GetLatestLoggerDataResponse); rpc ExperimentCommand (TrainerCommand) returns (CommandResponse); @@ -17,11 +17,31 @@ service ExperimentService { // Data Service (for weights_studio UI) rpc ApplyDataQuery (DataQueryRequest) returns (DataQueryResponse); rpc GetDataSamples (DataSamplesRequest) returns (DataSamplesResponse); - rpc EditDataSample (DataEditsRequest) returns (DataEditsResponse); + rpc EditDataSample (DataEditsRequest) returns (DataEditsRequest); rpc GetDataSplits (Empty) returns (DataSplitsResponse); rpc CheckAgentHealth (Empty) returns (AgentHealthResponse); + + // Checkpoint restore + rpc RestoreCheckpoint (RestoreCheckpointRequest) returns (RestoreCheckpointResponse); +} + +// --- Logger Data Sync --- +message GetLatestLoggerDataRequest { + bool request_full_history = 1; // true = full history, false = only queue (new data) + int32 max_points = 2; // max points per signal (only used for full history) +} + +message LoggerDataPoint { + string metric_name = 1; // The metric/signal name (e.g., "train/loss", "eval/accuracy") + int32 model_age = 2; + float metric_value = 3; + string experiment_hash = 4; + int64 timestamp = 5; } +message GetLatestLoggerDataResponse { + repeated LoggerDataPoint points = 1; // All points from all metrics +} message Empty {} @@ -347,3 +367,13 @@ message AgentHealthResponse { string message = 2; } +// --- Checkpoint Restore --- +message RestoreCheckpointRequest { + string experiment_hash = 1; // Hash code of checkpoint to restore +} + +message RestoreCheckpointResponse { + bool success = 1; // Whether restore was successful + string message = 2; // Details/error message +} + diff --git a/weightslab/proto/experiment_service_pb2.py b/weightslab/proto/experiment_service_pb2.py index e075ea15..c456e786 100644 --- a/weightslab/proto/experiment_service_pb2.py +++ b/weightslab/proto/experiment_service_pb2.py @@ -1,7 +1,7 @@ # -*- coding: utf-8 -*- # Generated by the protocol buffer compiler. DO NOT EDIT! # NO CHECKED-IN PROTOBUF GENCODE -# source: experiment_service.proto +# source: weightslab/proto/experiment_service.proto # Protobuf Python Version: 6.31.1 """Generated protocol buffer code.""" from google.protobuf import descriptor as _descriptor @@ -15,7 +15,7 @@ 31, 1, '', - 'experiment_service.proto' + 'weightslab/proto/experiment_service.proto' ) # @@protoc_insertion_point(imports) @@ -24,11 +24,11 @@ -DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x18\x65xperiment_service.proto\"\x07\n\x05\x45mpty\"/\n\x08NeuronId\x12\x10\n\x08layer_id\x18\x01 \x01(\x05\x12\x11\n\tneuron_id\x18\x02 \x01(\x05\"\x91\x02\n\x0fWeightOperation\x12*\n\x07op_type\x18\x01 \x01(\x0e\x32\x14.WeightOperationTypeH\x00\x88\x01\x01\x12\x15\n\x08layer_id\x18\x02 \x01(\x05H\x01\x88\x01\x01\x12\x1d\n\nneuron_ids\x18\x03 \x03(\x0b\x32\t.NeuronId\x12\x16\n\x0eneurons_to_add\x18\t \x01(\x05\x12 \n\x18zerofy_from_incoming_ids\x18\x0b \x03(\x05\x12\x1c\n\x14zerofy_to_neuron_ids\x18\x0c \x03(\x05\x12+\n\x11zerofy_predicates\x18\r \x03(\x0e\x32\x10.ZerofyPredicateB\n\n\x08_op_typeB\x0b\n\t_layer_id\"_\n\x17WeightsOperationRequest\x12/\n\x10weight_operation\x18\x01 \x01(\x0b\x32\x10.WeightOperationH\x00\x88\x01\x01\x42\x13\n\x11_weight_operation\"<\n\x18WeightsOperationResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"\xf3\x02\n\x0fHyperParameters\x12\x1c\n\x0f\x65xperiment_name\x18\x01 \x01(\tH\x00\x88\x01\x01\x12!\n\x14training_steps_to_do\x18\x02 \x01(\x05H\x01\x88\x01\x01\x12\x1a\n\rlearning_rate\x18\x03 \x01(\x02H\x02\x88\x01\x01\x12\x17\n\nbatch_size\x18\x04 \x01(\x05H\x03\x88\x01\x01\x12 \n\x13\x66ull_eval_frequency\x18\x05 \x01(\x05H\x04\x88\x01\x01\x12 \n\x13\x63heckpont_frequency\x18\x06 \x01(\x05H\x05\x88\x01\x01\x12\x18\n\x0bis_training\x18\x07 \x01(\x08H\x06\x88\x01\x01\x42\x12\n\x10_experiment_nameB\x17\n\x15_training_steps_to_doB\x10\n\x0e_learning_rateB\r\n\x0b_batch_sizeB\x16\n\x14_full_eval_frequencyB\x16\n\x14_checkpont_frequencyB\x0e\n\x0c_is_training\",\n\rMetricsStatus\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\x02\"~\n\rAnnotatStatus\x12\x0c\n\x04name\x18\x01 \x01(\t\x12.\n\x08metadata\x18\x02 \x03(\x0b\x32\x1c.AnnotatStatus.MetadataEntry\x1a/\n\rMetadataEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\x02:\x02\x38\x01\"\x90\x02\n\x10TrainingStatusEx\x12\x16\n\ttimestamp\x18\x01 \x01(\tH\x00\x88\x01\x01\x12\x1c\n\x0f\x65xperiment_name\x18\x02 \x01(\tH\x01\x88\x01\x01\x12\x16\n\tmodel_age\x18\x03 \x01(\x05H\x02\x88\x01\x01\x12+\n\x0emetrics_status\x18\x04 \x01(\x0b\x32\x0e.MetricsStatusH\x03\x88\x01\x01\x12+\n\x0e\x61nnotat_status\x18\x05 \x01(\x0b\x32\x0e.AnnotatStatusH\x04\x88\x01\x01\x42\x0c\n\n_timestampB\x12\n\x10_experiment_nameB\x0c\n\n_model_ageB\x11\n\x0f_metrics_statusB\x11\n\x0f_annotat_status\"]\n\x15HyperParameterCommand\x12/\n\x10hyper_parameters\x18\x01 \x01(\x0b\x32\x10.HyperParametersH\x00\x88\x01\x01\x42\x13\n\x11_hyper_parameters\">\n\x14\x44\x65nySamplesOperation\x12\x12\n\nsample_ids\x18\x01 \x03(\x05\x12\x12\n\naccumulate\x18\x02 \x01(\x08\"0\n\x17LoadCheckpointOperation\x12\x15\n\rcheckpoint_id\x18\x01 \x01(\x05\"\x8e\x06\n\x0eTrainerCommand\x12\x1c\n\x14get_hyper_parameters\x18\x04 \x01(\x08\x12\x1e\n\x16get_interactive_layers\x18\x05 \x01(\x08\x12\x1d\n\x10get_data_records\x18\x06 \x01(\tH\x00\x88\x01\x01\x12%\n\x18get_single_layer_info_id\x18\x08 \x01(\x05H\x01\x88\x01\x01\x12;\n\x16hyper_parameter_change\x18\x01 \x01(\x0b\x32\x16.HyperParameterCommandH\x02\x88\x01\x01\x12:\n\x16\x64\x65ny_samples_operation\x18\x07 \x01(\x0b\x32\x15.DenySamplesOperationH\x03\x88\x01\x01\x12?\n\x1b\x64\x65ny_eval_samples_operation\x18\n \x01(\x0b\x32\x15.DenySamplesOperationH\x04\x88\x01\x01\x12@\n\x19load_checkpoint_operation\x18\t \x01(\x0b\x32\x18.LoadCheckpointOperationH\x05\x88\x01\x01\x12\x42\n\x1eremove_from_denylist_operation\x18\x0b \x01(\x0b\x32\x15.DenySamplesOperationH\x06\x88\x01\x01\x12G\n#remove_eval_from_denylist_operation\x18\x0c \x01(\x0b\x32\x15.DenySamplesOperationH\x07\x88\x01\x01\x42\x13\n\x11_get_data_recordsB\x1b\n\x19_get_single_layer_info_idB\x19\n\x17_hyper_parameter_changeB\x19\n\x17_deny_samples_operationB\x1e\n\x1c_deny_eval_samples_operationB\x1c\n\x1a_load_checkpoint_operationB!\n\x1f_remove_from_denylist_operationB&\n$_remove_eval_from_denylist_operation\"\x9d\x01\n\x12HyperParameterDesc\x12\r\n\x05label\x18\x01 \x01(\t\x12\x0c\n\x04name\x18\x02 \x01(\t\x12\x0c\n\x04type\x18\x03 \x01(\t\x12\x1c\n\x0fnumerical_value\x18\x04 \x01(\x02H\x00\x88\x01\x01\x12\x19\n\x0cstring_value\x18\x05 \x01(\tH\x01\x88\x01\x01\x42\x12\n\x10_numerical_valueB\x0f\n\r_string_value\"\xf2\x02\n\x10NeuronStatistics\x12!\n\tneuron_id\x18\x01 \x01(\x0b\x32\t.NeuronIdH\x00\x88\x01\x01\x12\x17\n\nneuron_age\x18\x02 \x01(\x05H\x01\x88\x01\x01\x12\x1f\n\x12train_trigger_rate\x18\x03 \x01(\x02H\x02\x88\x01\x01\x12\x1e\n\x11\x65val_trigger_rate\x18\x04 \x01(\x02H\x03\x88\x01\x01\x12\x1a\n\rlearning_rate\x18\x07 \x01(\x02H\x04\x88\x01\x01\x12\x36\n\x0bincoming_lr\x18\x08 \x03(\x0b\x32!.NeuronStatistics.IncomingLrEntry\x1a\x31\n\x0fIncomingLrEntry\x12\x0b\n\x03key\x18\x01 \x01(\x05\x12\r\n\x05value\x18\x02 \x01(\x02:\x02\x38\x01\x42\x0c\n\n_neuron_idB\r\n\x0b_neuron_ageB\x15\n\x13_train_trigger_rateB\x14\n\x12_eval_trigger_rateB\x10\n\x0e_learning_rate\"\xf0\x02\n\x13LayerRepresentation\x12\x15\n\x08layer_id\x18\x01 \x01(\x05H\x00\x88\x01\x01\x12\x17\n\nlayer_name\x18\x02 \x01(\tH\x01\x88\x01\x01\x12\x17\n\nlayer_type\x18\x03 \x01(\tH\x02\x88\x01\x01\x12\x1a\n\rneurons_count\x18\x04 \x01(\x05H\x03\x88\x01\x01\x12#\n\x16incoming_neurons_count\x18\x05 \x01(\x05H\x04\x88\x01\x01\x12\x18\n\x0bkernel_size\x18\x06 \x01(\x05H\x05\x88\x01\x01\x12\x13\n\x06stride\x18\x07 \x01(\x05H\x06\x88\x01\x01\x12-\n\x12neurons_statistics\x18\n \x03(\x0b\x32\x11.NeuronStatisticsB\x0b\n\t_layer_idB\r\n\x0b_layer_nameB\r\n\x0b_layer_typeB\x10\n\x0e_neurons_countB\x19\n\x17_incoming_neurons_countB\x0e\n\x0c_kernel_sizeB\t\n\x07_stride\"H\n\x11\x41\x63tivationRequest\x12\x10\n\x08layer_id\x18\x01 \x01(\x05\x12\x11\n\tsample_id\x18\x02 \x01(\x05\x12\x0e\n\x06origin\x18\x03 \x01(\t\"H\n\rActivationMap\x12\x11\n\tneuron_id\x18\x01 \x01(\x05\x12\x0e\n\x06values\x18\x02 \x03(\x02\x12\t\n\x01H\x18\x03 \x01(\x05\x12\t\n\x01W\x18\x04 \x01(\x05\"d\n\x12\x41\x63tivationResponse\x12\x12\n\nlayer_type\x18\x01 \x01(\t\x12\x15\n\rneurons_count\x18\x02 \x01(\x05\x12#\n\x0b\x61\x63tivations\x18\x03 \x03(\x0b\x32\x0e.ActivationMap\"\x93\x01\n\tTaskField\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\x15\n\x0b\x66loat_value\x18\x02 \x01(\x02H\x00\x12\x13\n\tint_value\x18\x03 \x01(\x05H\x00\x12\x16\n\x0cstring_value\x18\x04 \x01(\tH\x00\x12\x15\n\x0b\x62ytes_value\x18\x05 \x01(\x0cH\x00\x12\x14\n\nbool_value\x18\x06 \x01(\x08H\x00\x42\x07\n\x05value\"\xcc\x02\n\x0eRecordMetadata\x12\x11\n\tsample_id\x18\x01 \x01(\x05\x12\x14\n\x0csample_label\x18\x02 \x03(\x05\x12\x19\n\x11sample_prediction\x18\x03 \x03(\x05\x12=\n\x10sample_last_loss\x18\x04 \x03(\x0b\x32#.RecordMetadata.SampleLastLossEntry\x12\x19\n\x11sample_encounters\x18\x05 \x01(\x05\x12\x18\n\x10sample_discarded\x18\x06 \x01(\x08\x12 \n\x0c\x65xtra_fields\x18\x07 \x03(\x0b\x32\n.TaskField\x12\x16\n\x0eprediction_raw\x18\t \x01(\x0c\x12\x11\n\ttask_type\x18\n \x01(\t\x1a\x35\n\x13SampleLastLossEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\x02:\x02\x38\x01\"\x93\x01\n\x10SampleStatistics\x12\x13\n\x06origin\x18\x06 \x01(\tH\x00\x88\x01\x01\x12\x19\n\x0csample_count\x18\x07 \x01(\x05H\x01\x88\x01\x01\x12\x11\n\ttask_type\x18\t \x01(\t\x12 \n\x07records\x18\x08 \x03(\x0b\x32\x0f.RecordMetadataB\t\n\x07_originB\x0f\n\r_sample_count\"\xe6\x01\n\x0f\x43ommandResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\x33\n\x16hyper_parameters_descs\x18\x03 \x03(\x0b\x32\x13.HyperParameterDesc\x12\x33\n\x15layer_representations\x18\x04 \x03(\x0b\x32\x14.LayerRepresentation\x12\x31\n\x11sample_statistics\x18\x05 \x01(\x0b\x32\x11.SampleStatisticsH\x00\x88\x01\x01\x42\x14\n\x12_sample_statistics\"U\n\rSampleRequest\x12\x16\n\tsample_id\x18\x01 \x01(\x05H\x00\x88\x01\x01\x12\x13\n\x06origin\x18\x02 \x01(\tH\x01\x88\x01\x01\x42\x0c\n\n_sample_idB\t\n\x07_origin\"\xad\x02\n\x15SampleRequestResponse\x12\x16\n\tsample_id\x18\x01 \x01(\x05H\x00\x88\x01\x01\x12\x13\n\x06origin\x18\x02 \x01(\tH\x01\x88\x01\x01\x12\x12\n\x05label\x18\x03 \x01(\x05H\x02\x88\x01\x01\x12\x11\n\x04\x64\x61ta\x18\x04 \x01(\x0cH\x03\x88\x01\x01\x12\x1a\n\rerror_message\x18\x05 \x01(\tH\x04\x88\x01\x01\x12\x15\n\x08raw_data\x18\x06 \x01(\x0cH\x05\x88\x01\x01\x12\x11\n\x04mask\x18\x07 \x01(\x0cH\x06\x88\x01\x01\x12\x17\n\nprediction\x18\x08 \x01(\x0cH\x07\x88\x01\x01\x42\x0c\n\n_sample_idB\t\n\x07_originB\x08\n\x06_labelB\x07\n\x05_dataB\x10\n\x0e_error_messageB\x0b\n\t_raw_dataB\x07\n\x05_maskB\r\n\x0b_prediction\"\x92\x01\n\x12\x42\x61tchSampleRequest\x12\x12\n\nsample_ids\x18\x01 \x03(\x05\x12\x0e\n\x06origin\x18\x02 \x01(\t\x12\x19\n\x0cresize_width\x18\x03 \x01(\x05H\x00\x88\x01\x01\x12\x1a\n\rresize_height\x18\x04 \x01(\x05H\x01\x88\x01\x01\x42\x0f\n\r_resize_widthB\x10\n\x0e_resize_height\">\n\x13\x42\x61tchSampleResponse\x12\'\n\x07samples\x18\x01 \x03(\x0b\x32\x16.SampleRequestResponse\".\n\x0eWeightsRequest\x12\x1c\n\tneuron_id\x18\x01 \x01(\x0b\x32\t.NeuronId\"\x9d\x02\n\x0fWeightsResponse\x12\x1c\n\tneuron_id\x18\x01 \x01(\x0b\x32\t.NeuronId\x12\x17\n\nlayer_name\x18\x02 \x01(\tH\x00\x88\x01\x01\x12\x17\n\nlayer_type\x18\x03 \x01(\tH\x01\x88\x01\x01\x12\x10\n\x08incoming\x18\x04 \x01(\x05\x12\x10\n\x08outgoing\x18\x05 \x01(\x05\x12\x18\n\x0bkernel_size\x18\x06 \x01(\x05H\x02\x88\x01\x01\x12\x0f\n\x07weights\x18\x07 \x03(\x02\x12\x0f\n\x07success\x18\x0b \x01(\x08\x12\x1a\n\rerror_message\x18\x0c \x01(\tH\x03\x88\x01\x01\x42\r\n\x0b_layer_nameB\r\n\x0b_layer_typeB\x0e\n\x0c_kernel_sizeB\x10\n\x0e_error_message\"R\n\x10\x44\x61taQueryRequest\x12\r\n\x05query\x18\x01 \x01(\t\x12\x12\n\naccumulate\x18\x02 \x01(\x08\x12\x1b\n\x13is_natural_language\x18\x03 \x01(\x08\"\xfb\x01\n\x11\x44\x61taQueryResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\x1d\n\x15number_of_all_samples\x18\x03 \x01(\x05\x12%\n\x1dnumber_of_samples_in_the_loop\x18\x04 \x01(\x05\x12#\n\x1bnumber_of_discarded_samples\x18\x05 \x01(\x05\x12\x13\n\x0bunique_tags\x18\x06 \x03(\t\x12+\n\x11\x61gent_intent_type\x18\x07 \x01(\x0e\x32\x10.AgentIntentType\x12\x17\n\x0f\x61nalysis_result\x18\x08 \x01(\t\"\xc2\x01\n\x12\x44\x61taSamplesRequest\x12\x13\n\x0bstart_index\x18\x01 \x01(\x05\x12\x13\n\x0brecords_cnt\x18\x02 \x01(\x05\x12 \n\x18include_transformed_data\x18\x03 \x01(\x08\x12\x18\n\x10include_raw_data\x18\x04 \x01(\x08\x12\x19\n\x11stats_to_retrieve\x18\x05 \x03(\t\x12\x14\n\x0cresize_width\x18\x06 \x01(\x05\x12\x15\n\rresize_height\x18\x07 \x01(\x05\"m\n\x08\x44\x61taStat\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\x0c\n\x04type\x18\x02 \x01(\t\x12\r\n\x05shape\x18\x03 \x03(\x05\x12\r\n\x05value\x18\x04 \x03(\x02\x12\x14\n\x0cvalue_string\x18\x05 \x01(\t\x12\x11\n\tthumbnail\x18\x06 \x01(\x0c\">\n\nDataRecord\x12\x11\n\tsample_id\x18\x01 \x01(\x05\x12\x1d\n\ndata_stats\x18\x02 \x03(\x0b\x32\t.DataStat\"Z\n\x13\x44\x61taSamplesResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12!\n\x0c\x64\x61ta_records\x18\x03 \x03(\x0b\x32\x0b.DataRecord\"\xb0\x01\n\x10\x44\x61taEditsRequest\x12\x11\n\tstat_name\x18\x01 \x01(\t\x12\x13\n\x0b\x66loat_value\x18\x02 \x01(\x02\x12\x14\n\x0cstring_value\x18\x03 \x01(\t\x12\x12\n\nbool_value\x18\x04 \x01(\x08\x12\x1d\n\x04type\x18\x05 \x01(\x0e\x32\x0f.SampleEditType\x12\x13\n\x0bsamples_ids\x18\x06 \x03(\x05\x12\x16\n\x0esample_origins\x18\x07 \x03(\t\"5\n\x11\x44\x61taEditsResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\":\n\x12\x44\x61taSplitsResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x13\n\x0bsplit_names\x18\x02 \x03(\t\"9\n\x13\x41gentHealthResponse\x12\x11\n\tavailable\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t*d\n\x13WeightOperationType\x12\n\n\x06ZEROFY\x10\x00\x12\x10\n\x0cREINITIALIZE\x10\x01\x12\n\n\x06\x46REEZE\x10\x02\x12\x12\n\x0eREMOVE_NEURONS\x10\t\x12\x0f\n\x0b\x41\x44\x44_NEURONS\x10\n*o\n\x0fZerofyPredicate\x12\x19\n\x15ZEROFY_PREDICATE_NONE\x10\x00\x12 \n\x1cZEROFY_PREDICATE_WITH_FROZEN\x10\x01\x12\x1f\n\x1bZEROFY_PREDICATE_WITH_OLDER\x10\x02*M\n\x0f\x41gentIntentType\x12\x12\n\x0eINTENT_UNKNOWN\x10\x00\x12\x11\n\rINTENT_FILTER\x10\x01\x12\x13\n\x0fINTENT_ANALYSIS\x10\x02*I\n\x0eSampleEditType\x12\x11\n\rEDIT_OVERRIDE\x10\x00\x12\x13\n\x0f\x45\x44IT_ACCUMULATE\x10\x01\x12\x0f\n\x0b\x45\x44IT_REMOVE\x10\x02\x32\xf6\x04\n\x11\x45xperimentService\x12+\n\x0cStreamStatus\x12\x06.Empty\x1a\x11.TrainingStatusEx0\x01\x12\x36\n\x11\x45xperimentCommand\x12\x0f.TrainerCommand\x1a\x10.CommandResponse\x12H\n\x11ManipulateWeights\x12\x18.WeightsOperationRequest\x1a\x19.WeightsOperationResponse\x12/\n\nGetWeights\x12\x0f.WeightsRequest\x1a\x10.WeightsResponse\x12\x39\n\x0eGetActivations\x12\x12.ActivationRequest\x1a\x13.ActivationResponse\x12\x37\n\nGetSamples\x12\x13.BatchSampleRequest\x1a\x14.BatchSampleResponse\x12\x37\n\x0e\x41pplyDataQuery\x12\x11.DataQueryRequest\x1a\x12.DataQueryResponse\x12;\n\x0eGetDataSamples\x12\x13.DataSamplesRequest\x1a\x14.DataSamplesResponse\x12\x37\n\x0e\x45\x64itDataSample\x12\x11.DataEditsRequest\x1a\x12.DataEditsResponse\x12,\n\rGetDataSplits\x12\x06.Empty\x1a\x13.DataSplitsResponse\x12\x30\n\x10\x43heckAgentHealth\x12\x06.Empty\x1a\x14.AgentHealthResponseb\x06proto3') +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n)weightslab/proto/experiment_service.proto\"N\n\x1aGetLatestLoggerDataRequest\x12\x1c\n\x14request_full_history\x18\x01 \x01(\x08\x12\x12\n\nmax_points\x18\x02 \x01(\x05\"{\n\x0fLoggerDataPoint\x12\x13\n\x0bmetric_name\x18\x01 \x01(\t\x12\x11\n\tmodel_age\x18\x02 \x01(\x05\x12\x14\n\x0cmetric_value\x18\x03 \x01(\x02\x12\x17\n\x0f\x65xperiment_hash\x18\x04 \x01(\t\x12\x11\n\ttimestamp\x18\x05 \x01(\x03\"?\n\x1bGetLatestLoggerDataResponse\x12 \n\x06points\x18\x01 \x03(\x0b\x32\x10.LoggerDataPoint\"\x07\n\x05\x45mpty\"/\n\x08NeuronId\x12\x10\n\x08layer_id\x18\x01 \x01(\x05\x12\x11\n\tneuron_id\x18\x02 \x01(\x05\"\x91\x02\n\x0fWeightOperation\x12*\n\x07op_type\x18\x01 \x01(\x0e\x32\x14.WeightOperationTypeH\x00\x88\x01\x01\x12\x15\n\x08layer_id\x18\x02 \x01(\x05H\x01\x88\x01\x01\x12\x1d\n\nneuron_ids\x18\x03 \x03(\x0b\x32\t.NeuronId\x12\x16\n\x0eneurons_to_add\x18\t \x01(\x05\x12 \n\x18zerofy_from_incoming_ids\x18\x0b \x03(\x05\x12\x1c\n\x14zerofy_to_neuron_ids\x18\x0c \x03(\x05\x12+\n\x11zerofy_predicates\x18\r \x03(\x0e\x32\x10.ZerofyPredicateB\n\n\x08_op_typeB\x0b\n\t_layer_id\"_\n\x17WeightsOperationRequest\x12/\n\x10weight_operation\x18\x01 \x01(\x0b\x32\x10.WeightOperationH\x00\x88\x01\x01\x42\x13\n\x11_weight_operation\"<\n\x18WeightsOperationResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"\xf3\x02\n\x0fHyperParameters\x12\x1c\n\x0f\x65xperiment_name\x18\x01 \x01(\tH\x00\x88\x01\x01\x12!\n\x14training_steps_to_do\x18\x02 \x01(\x05H\x01\x88\x01\x01\x12\x1a\n\rlearning_rate\x18\x03 \x01(\x02H\x02\x88\x01\x01\x12\x17\n\nbatch_size\x18\x04 \x01(\x05H\x03\x88\x01\x01\x12 \n\x13\x66ull_eval_frequency\x18\x05 \x01(\x05H\x04\x88\x01\x01\x12 \n\x13\x63heckpont_frequency\x18\x06 \x01(\x05H\x05\x88\x01\x01\x12\x18\n\x0bis_training\x18\x07 \x01(\x08H\x06\x88\x01\x01\x42\x12\n\x10_experiment_nameB\x17\n\x15_training_steps_to_doB\x10\n\x0e_learning_rateB\r\n\x0b_batch_sizeB\x16\n\x14_full_eval_frequencyB\x16\n\x14_checkpont_frequencyB\x0e\n\x0c_is_training\",\n\rMetricsStatus\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\x02\"~\n\rAnnotatStatus\x12\x0c\n\x04name\x18\x01 \x01(\t\x12.\n\x08metadata\x18\x02 \x03(\x0b\x32\x1c.AnnotatStatus.MetadataEntry\x1a/\n\rMetadataEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\x02:\x02\x38\x01\"\x90\x02\n\x10TrainingStatusEx\x12\x16\n\ttimestamp\x18\x01 \x01(\tH\x00\x88\x01\x01\x12\x1c\n\x0f\x65xperiment_name\x18\x02 \x01(\tH\x01\x88\x01\x01\x12\x16\n\tmodel_age\x18\x03 \x01(\x05H\x02\x88\x01\x01\x12+\n\x0emetrics_status\x18\x04 \x01(\x0b\x32\x0e.MetricsStatusH\x03\x88\x01\x01\x12+\n\x0e\x61nnotat_status\x18\x05 \x01(\x0b\x32\x0e.AnnotatStatusH\x04\x88\x01\x01\x42\x0c\n\n_timestampB\x12\n\x10_experiment_nameB\x0c\n\n_model_ageB\x11\n\x0f_metrics_statusB\x11\n\x0f_annotat_status\"]\n\x15HyperParameterCommand\x12/\n\x10hyper_parameters\x18\x01 \x01(\x0b\x32\x10.HyperParametersH\x00\x88\x01\x01\x42\x13\n\x11_hyper_parameters\">\n\x14\x44\x65nySamplesOperation\x12\x12\n\nsample_ids\x18\x01 \x03(\x05\x12\x12\n\naccumulate\x18\x02 \x01(\x08\"0\n\x17LoadCheckpointOperation\x12\x15\n\rcheckpoint_id\x18\x01 \x01(\x05\"\x8e\x06\n\x0eTrainerCommand\x12\x1c\n\x14get_hyper_parameters\x18\x04 \x01(\x08\x12\x1e\n\x16get_interactive_layers\x18\x05 \x01(\x08\x12\x1d\n\x10get_data_records\x18\x06 \x01(\tH\x00\x88\x01\x01\x12%\n\x18get_single_layer_info_id\x18\x08 \x01(\x05H\x01\x88\x01\x01\x12;\n\x16hyper_parameter_change\x18\x01 \x01(\x0b\x32\x16.HyperParameterCommandH\x02\x88\x01\x01\x12:\n\x16\x64\x65ny_samples_operation\x18\x07 \x01(\x0b\x32\x15.DenySamplesOperationH\x03\x88\x01\x01\x12?\n\x1b\x64\x65ny_eval_samples_operation\x18\n \x01(\x0b\x32\x15.DenySamplesOperationH\x04\x88\x01\x01\x12@\n\x19load_checkpoint_operation\x18\t \x01(\x0b\x32\x18.LoadCheckpointOperationH\x05\x88\x01\x01\x12\x42\n\x1eremove_from_denylist_operation\x18\x0b \x01(\x0b\x32\x15.DenySamplesOperationH\x06\x88\x01\x01\x12G\n#remove_eval_from_denylist_operation\x18\x0c \x01(\x0b\x32\x15.DenySamplesOperationH\x07\x88\x01\x01\x42\x13\n\x11_get_data_recordsB\x1b\n\x19_get_single_layer_info_idB\x19\n\x17_hyper_parameter_changeB\x19\n\x17_deny_samples_operationB\x1e\n\x1c_deny_eval_samples_operationB\x1c\n\x1a_load_checkpoint_operationB!\n\x1f_remove_from_denylist_operationB&\n$_remove_eval_from_denylist_operation\"\x9d\x01\n\x12HyperParameterDesc\x12\r\n\x05label\x18\x01 \x01(\t\x12\x0c\n\x04name\x18\x02 \x01(\t\x12\x0c\n\x04type\x18\x03 \x01(\t\x12\x1c\n\x0fnumerical_value\x18\x04 \x01(\x02H\x00\x88\x01\x01\x12\x19\n\x0cstring_value\x18\x05 \x01(\tH\x01\x88\x01\x01\x42\x12\n\x10_numerical_valueB\x0f\n\r_string_value\"\xf2\x02\n\x10NeuronStatistics\x12!\n\tneuron_id\x18\x01 \x01(\x0b\x32\t.NeuronIdH\x00\x88\x01\x01\x12\x17\n\nneuron_age\x18\x02 \x01(\x05H\x01\x88\x01\x01\x12\x1f\n\x12train_trigger_rate\x18\x03 \x01(\x02H\x02\x88\x01\x01\x12\x1e\n\x11\x65val_trigger_rate\x18\x04 \x01(\x02H\x03\x88\x01\x01\x12\x1a\n\rlearning_rate\x18\x07 \x01(\x02H\x04\x88\x01\x01\x12\x36\n\x0bincoming_lr\x18\x08 \x03(\x0b\x32!.NeuronStatistics.IncomingLrEntry\x1a\x31\n\x0fIncomingLrEntry\x12\x0b\n\x03key\x18\x01 \x01(\x05\x12\r\n\x05value\x18\x02 \x01(\x02:\x02\x38\x01\x42\x0c\n\n_neuron_idB\r\n\x0b_neuron_ageB\x15\n\x13_train_trigger_rateB\x14\n\x12_eval_trigger_rateB\x10\n\x0e_learning_rate\"\xf0\x02\n\x13LayerRepresentation\x12\x15\n\x08layer_id\x18\x01 \x01(\x05H\x00\x88\x01\x01\x12\x17\n\nlayer_name\x18\x02 \x01(\tH\x01\x88\x01\x01\x12\x17\n\nlayer_type\x18\x03 \x01(\tH\x02\x88\x01\x01\x12\x1a\n\rneurons_count\x18\x04 \x01(\x05H\x03\x88\x01\x01\x12#\n\x16incoming_neurons_count\x18\x05 \x01(\x05H\x04\x88\x01\x01\x12\x18\n\x0bkernel_size\x18\x06 \x01(\x05H\x05\x88\x01\x01\x12\x13\n\x06stride\x18\x07 \x01(\x05H\x06\x88\x01\x01\x12-\n\x12neurons_statistics\x18\n \x03(\x0b\x32\x11.NeuronStatisticsB\x0b\n\t_layer_idB\r\n\x0b_layer_nameB\r\n\x0b_layer_typeB\x10\n\x0e_neurons_countB\x19\n\x17_incoming_neurons_countB\x0e\n\x0c_kernel_sizeB\t\n\x07_stride\"H\n\x11\x41\x63tivationRequest\x12\x10\n\x08layer_id\x18\x01 \x01(\x05\x12\x11\n\tsample_id\x18\x02 \x01(\x05\x12\x0e\n\x06origin\x18\x03 \x01(\t\"H\n\rActivationMap\x12\x11\n\tneuron_id\x18\x01 \x01(\x05\x12\x0e\n\x06values\x18\x02 \x03(\x02\x12\t\n\x01H\x18\x03 \x01(\x05\x12\t\n\x01W\x18\x04 \x01(\x05\"d\n\x12\x41\x63tivationResponse\x12\x12\n\nlayer_type\x18\x01 \x01(\t\x12\x15\n\rneurons_count\x18\x02 \x01(\x05\x12#\n\x0b\x61\x63tivations\x18\x03 \x03(\x0b\x32\x0e.ActivationMap\"\x93\x01\n\tTaskField\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\x15\n\x0b\x66loat_value\x18\x02 \x01(\x02H\x00\x12\x13\n\tint_value\x18\x03 \x01(\x05H\x00\x12\x16\n\x0cstring_value\x18\x04 \x01(\tH\x00\x12\x15\n\x0b\x62ytes_value\x18\x05 \x01(\x0cH\x00\x12\x14\n\nbool_value\x18\x06 \x01(\x08H\x00\x42\x07\n\x05value\"\xcc\x02\n\x0eRecordMetadata\x12\x11\n\tsample_id\x18\x01 \x01(\x05\x12\x14\n\x0csample_label\x18\x02 \x03(\x05\x12\x19\n\x11sample_prediction\x18\x03 \x03(\x05\x12=\n\x10sample_last_loss\x18\x04 \x03(\x0b\x32#.RecordMetadata.SampleLastLossEntry\x12\x19\n\x11sample_encounters\x18\x05 \x01(\x05\x12\x18\n\x10sample_discarded\x18\x06 \x01(\x08\x12 \n\x0c\x65xtra_fields\x18\x07 \x03(\x0b\x32\n.TaskField\x12\x16\n\x0eprediction_raw\x18\t \x01(\x0c\x12\x11\n\ttask_type\x18\n \x01(\t\x1a\x35\n\x13SampleLastLossEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\x02:\x02\x38\x01\"\x93\x01\n\x10SampleStatistics\x12\x13\n\x06origin\x18\x06 \x01(\tH\x00\x88\x01\x01\x12\x19\n\x0csample_count\x18\x07 \x01(\x05H\x01\x88\x01\x01\x12\x11\n\ttask_type\x18\t \x01(\t\x12 \n\x07records\x18\x08 \x03(\x0b\x32\x0f.RecordMetadataB\t\n\x07_originB\x0f\n\r_sample_count\"\xe6\x01\n\x0f\x43ommandResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\x33\n\x16hyper_parameters_descs\x18\x03 \x03(\x0b\x32\x13.HyperParameterDesc\x12\x33\n\x15layer_representations\x18\x04 \x03(\x0b\x32\x14.LayerRepresentation\x12\x31\n\x11sample_statistics\x18\x05 \x01(\x0b\x32\x11.SampleStatisticsH\x00\x88\x01\x01\x42\x14\n\x12_sample_statistics\"U\n\rSampleRequest\x12\x16\n\tsample_id\x18\x01 \x01(\x05H\x00\x88\x01\x01\x12\x13\n\x06origin\x18\x02 \x01(\tH\x01\x88\x01\x01\x42\x0c\n\n_sample_idB\t\n\x07_origin\"\xad\x02\n\x15SampleRequestResponse\x12\x16\n\tsample_id\x18\x01 \x01(\x05H\x00\x88\x01\x01\x12\x13\n\x06origin\x18\x02 \x01(\tH\x01\x88\x01\x01\x12\x12\n\x05label\x18\x03 \x01(\x05H\x02\x88\x01\x01\x12\x11\n\x04\x64\x61ta\x18\x04 \x01(\x0cH\x03\x88\x01\x01\x12\x1a\n\rerror_message\x18\x05 \x01(\tH\x04\x88\x01\x01\x12\x15\n\x08raw_data\x18\x06 \x01(\x0cH\x05\x88\x01\x01\x12\x11\n\x04mask\x18\x07 \x01(\x0cH\x06\x88\x01\x01\x12\x17\n\nprediction\x18\x08 \x01(\x0cH\x07\x88\x01\x01\x42\x0c\n\n_sample_idB\t\n\x07_originB\x08\n\x06_labelB\x07\n\x05_dataB\x10\n\x0e_error_messageB\x0b\n\t_raw_dataB\x07\n\x05_maskB\r\n\x0b_prediction\"\x92\x01\n\x12\x42\x61tchSampleRequest\x12\x12\n\nsample_ids\x18\x01 \x03(\x05\x12\x0e\n\x06origin\x18\x02 \x01(\t\x12\x19\n\x0cresize_width\x18\x03 \x01(\x05H\x00\x88\x01\x01\x12\x1a\n\rresize_height\x18\x04 \x01(\x05H\x01\x88\x01\x01\x42\x0f\n\r_resize_widthB\x10\n\x0e_resize_height\">\n\x13\x42\x61tchSampleResponse\x12\'\n\x07samples\x18\x01 \x03(\x0b\x32\x16.SampleRequestResponse\".\n\x0eWeightsRequest\x12\x1c\n\tneuron_id\x18\x01 \x01(\x0b\x32\t.NeuronId\"\x9d\x02\n\x0fWeightsResponse\x12\x1c\n\tneuron_id\x18\x01 \x01(\x0b\x32\t.NeuronId\x12\x17\n\nlayer_name\x18\x02 \x01(\tH\x00\x88\x01\x01\x12\x17\n\nlayer_type\x18\x03 \x01(\tH\x01\x88\x01\x01\x12\x10\n\x08incoming\x18\x04 \x01(\x05\x12\x10\n\x08outgoing\x18\x05 \x01(\x05\x12\x18\n\x0bkernel_size\x18\x06 \x01(\x05H\x02\x88\x01\x01\x12\x0f\n\x07weights\x18\x07 \x03(\x02\x12\x0f\n\x07success\x18\x0b \x01(\x08\x12\x1a\n\rerror_message\x18\x0c \x01(\tH\x03\x88\x01\x01\x42\r\n\x0b_layer_nameB\r\n\x0b_layer_typeB\x0e\n\x0c_kernel_sizeB\x10\n\x0e_error_message\"R\n\x10\x44\x61taQueryRequest\x12\r\n\x05query\x18\x01 \x01(\t\x12\x12\n\naccumulate\x18\x02 \x01(\x08\x12\x1b\n\x13is_natural_language\x18\x03 \x01(\x08\"\xfb\x01\n\x11\x44\x61taQueryResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12\x1d\n\x15number_of_all_samples\x18\x03 \x01(\x05\x12%\n\x1dnumber_of_samples_in_the_loop\x18\x04 \x01(\x05\x12#\n\x1bnumber_of_discarded_samples\x18\x05 \x01(\x05\x12\x13\n\x0bunique_tags\x18\x06 \x03(\t\x12+\n\x11\x61gent_intent_type\x18\x07 \x01(\x0e\x32\x10.AgentIntentType\x12\x17\n\x0f\x61nalysis_result\x18\x08 \x01(\t\"\xc2\x01\n\x12\x44\x61taSamplesRequest\x12\x13\n\x0bstart_index\x18\x01 \x01(\x05\x12\x13\n\x0brecords_cnt\x18\x02 \x01(\x05\x12 \n\x18include_transformed_data\x18\x03 \x01(\x08\x12\x18\n\x10include_raw_data\x18\x04 \x01(\x08\x12\x19\n\x11stats_to_retrieve\x18\x05 \x03(\t\x12\x14\n\x0cresize_width\x18\x06 \x01(\x05\x12\x15\n\rresize_height\x18\x07 \x01(\x05\"m\n\x08\x44\x61taStat\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\x0c\n\x04type\x18\x02 \x01(\t\x12\r\n\x05shape\x18\x03 \x03(\x05\x12\r\n\x05value\x18\x04 \x03(\x02\x12\x14\n\x0cvalue_string\x18\x05 \x01(\t\x12\x11\n\tthumbnail\x18\x06 \x01(\x0c\">\n\nDataRecord\x12\x11\n\tsample_id\x18\x01 \x01(\x05\x12\x1d\n\ndata_stats\x18\x02 \x03(\x0b\x32\t.DataStat\"Z\n\x13\x44\x61taSamplesResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\x12!\n\x0c\x64\x61ta_records\x18\x03 \x03(\x0b\x32\x0b.DataRecord\"\xb0\x01\n\x10\x44\x61taEditsRequest\x12\x11\n\tstat_name\x18\x01 \x01(\t\x12\x13\n\x0b\x66loat_value\x18\x02 \x01(\x02\x12\x14\n\x0cstring_value\x18\x03 \x01(\t\x12\x12\n\nbool_value\x18\x04 \x01(\x08\x12\x1d\n\x04type\x18\x05 \x01(\x0e\x32\x0f.SampleEditType\x12\x13\n\x0bsamples_ids\x18\x06 \x03(\x05\x12\x16\n\x0esample_origins\x18\x07 \x03(\t\"5\n\x11\x44\x61taEditsResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\":\n\x12\x44\x61taSplitsResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x13\n\x0bsplit_names\x18\x02 \x03(\t\"9\n\x13\x41gentHealthResponse\x12\x11\n\tavailable\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"3\n\x18RestoreCheckpointRequest\x12\x17\n\x0f\x65xperiment_hash\x18\x01 \x01(\t\"=\n\x19RestoreCheckpointResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t*d\n\x13WeightOperationType\x12\n\n\x06ZEROFY\x10\x00\x12\x10\n\x0cREINITIALIZE\x10\x01\x12\n\n\x06\x46REEZE\x10\x02\x12\x12\n\x0eREMOVE_NEURONS\x10\t\x12\x0f\n\x0b\x41\x44\x44_NEURONS\x10\n*o\n\x0fZerofyPredicate\x12\x19\n\x15ZEROFY_PREDICATE_NONE\x10\x00\x12 \n\x1cZEROFY_PREDICATE_WITH_FROZEN\x10\x01\x12\x1f\n\x1bZEROFY_PREDICATE_WITH_OLDER\x10\x02*M\n\x0f\x41gentIntentType\x12\x12\n\x0eINTENT_UNKNOWN\x10\x00\x12\x11\n\rINTENT_FILTER\x10\x01\x12\x13\n\x0fINTENT_ANALYSIS\x10\x02*I\n\x0eSampleEditType\x12\x11\n\rEDIT_OVERRIDE\x10\x00\x12\x13\n\x0f\x45\x44IT_ACCUMULATE\x10\x01\x12\x0f\n\x0b\x45\x44IT_REMOVE\x10\x02\x32\xe6\x05\n\x11\x45xperimentService\x12P\n\x13GetLatestLoggerData\x12\x1b.GetLatestLoggerDataRequest\x1a\x1c.GetLatestLoggerDataResponse\x12\x36\n\x11\x45xperimentCommand\x12\x0f.TrainerCommand\x1a\x10.CommandResponse\x12H\n\x11ManipulateWeights\x12\x18.WeightsOperationRequest\x1a\x19.WeightsOperationResponse\x12/\n\nGetWeights\x12\x0f.WeightsRequest\x1a\x10.WeightsResponse\x12\x39\n\x0eGetActivations\x12\x12.ActivationRequest\x1a\x13.ActivationResponse\x12\x37\n\nGetSamples\x12\x13.BatchSampleRequest\x1a\x14.BatchSampleResponse\x12\x37\n\x0e\x41pplyDataQuery\x12\x11.DataQueryRequest\x1a\x12.DataQueryResponse\x12;\n\x0eGetDataSamples\x12\x13.DataSamplesRequest\x1a\x14.DataSamplesResponse\x12\x36\n\x0e\x45\x64itDataSample\x12\x11.DataEditsRequest\x1a\x11.DataEditsRequest\x12,\n\rGetDataSplits\x12\x06.Empty\x1a\x13.DataSplitsResponse\x12\x30\n\x10\x43heckAgentHealth\x12\x06.Empty\x1a\x14.AgentHealthResponse\x12J\n\x11RestoreCheckpoint\x12\x19.RestoreCheckpointRequest\x1a\x1a.RestoreCheckpointResponseb\x06proto3') _globals = globals() _builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) -_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, 'experiment_service_pb2', _globals) +_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, 'weightslab.proto.experiment_service_pb2', _globals) if not _descriptor._USE_C_DESCRIPTORS: DESCRIPTOR._loaded_options = None _globals['_ANNOTATSTATUS_METADATAENTRY']._loaded_options = None @@ -37,98 +37,108 @@ _globals['_NEURONSTATISTICS_INCOMINGLRENTRY']._serialized_options = b'8\001' _globals['_RECORDMETADATA_SAMPLELASTLOSSENTRY']._loaded_options = None _globals['_RECORDMETADATA_SAMPLELASTLOSSENTRY']._serialized_options = b'8\001' - _globals['_WEIGHTOPERATIONTYPE']._serialized_start=6455 - _globals['_WEIGHTOPERATIONTYPE']._serialized_end=6555 - _globals['_ZEROFYPREDICATE']._serialized_start=6557 - _globals['_ZEROFYPREDICATE']._serialized_end=6668 - _globals['_AGENTINTENTTYPE']._serialized_start=6670 - _globals['_AGENTINTENTTYPE']._serialized_end=6747 - _globals['_SAMPLEEDITTYPE']._serialized_start=6749 - _globals['_SAMPLEEDITTYPE']._serialized_end=6822 - _globals['_EMPTY']._serialized_start=28 - _globals['_EMPTY']._serialized_end=35 - _globals['_NEURONID']._serialized_start=37 - _globals['_NEURONID']._serialized_end=84 - _globals['_WEIGHTOPERATION']._serialized_start=87 - _globals['_WEIGHTOPERATION']._serialized_end=360 - _globals['_WEIGHTSOPERATIONREQUEST']._serialized_start=362 - _globals['_WEIGHTSOPERATIONREQUEST']._serialized_end=457 - _globals['_WEIGHTSOPERATIONRESPONSE']._serialized_start=459 - _globals['_WEIGHTSOPERATIONRESPONSE']._serialized_end=519 - _globals['_HYPERPARAMETERS']._serialized_start=522 - _globals['_HYPERPARAMETERS']._serialized_end=893 - _globals['_METRICSSTATUS']._serialized_start=895 - _globals['_METRICSSTATUS']._serialized_end=939 - _globals['_ANNOTATSTATUS']._serialized_start=941 - _globals['_ANNOTATSTATUS']._serialized_end=1067 - _globals['_ANNOTATSTATUS_METADATAENTRY']._serialized_start=1020 - _globals['_ANNOTATSTATUS_METADATAENTRY']._serialized_end=1067 - _globals['_TRAININGSTATUSEX']._serialized_start=1070 - _globals['_TRAININGSTATUSEX']._serialized_end=1342 - _globals['_HYPERPARAMETERCOMMAND']._serialized_start=1344 - _globals['_HYPERPARAMETERCOMMAND']._serialized_end=1437 - _globals['_DENYSAMPLESOPERATION']._serialized_start=1439 - _globals['_DENYSAMPLESOPERATION']._serialized_end=1501 - _globals['_LOADCHECKPOINTOPERATION']._serialized_start=1503 - _globals['_LOADCHECKPOINTOPERATION']._serialized_end=1551 - _globals['_TRAINERCOMMAND']._serialized_start=1554 - _globals['_TRAINERCOMMAND']._serialized_end=2336 - _globals['_HYPERPARAMETERDESC']._serialized_start=2339 - _globals['_HYPERPARAMETERDESC']._serialized_end=2496 - _globals['_NEURONSTATISTICS']._serialized_start=2499 - _globals['_NEURONSTATISTICS']._serialized_end=2869 - _globals['_NEURONSTATISTICS_INCOMINGLRENTRY']._serialized_start=2728 - _globals['_NEURONSTATISTICS_INCOMINGLRENTRY']._serialized_end=2777 - _globals['_LAYERREPRESENTATION']._serialized_start=2872 - _globals['_LAYERREPRESENTATION']._serialized_end=3240 - _globals['_ACTIVATIONREQUEST']._serialized_start=3242 - _globals['_ACTIVATIONREQUEST']._serialized_end=3314 - _globals['_ACTIVATIONMAP']._serialized_start=3316 - _globals['_ACTIVATIONMAP']._serialized_end=3388 - _globals['_ACTIVATIONRESPONSE']._serialized_start=3390 - _globals['_ACTIVATIONRESPONSE']._serialized_end=3490 - _globals['_TASKFIELD']._serialized_start=3493 - _globals['_TASKFIELD']._serialized_end=3640 - _globals['_RECORDMETADATA']._serialized_start=3643 - _globals['_RECORDMETADATA']._serialized_end=3975 - _globals['_RECORDMETADATA_SAMPLELASTLOSSENTRY']._serialized_start=3922 - _globals['_RECORDMETADATA_SAMPLELASTLOSSENTRY']._serialized_end=3975 - _globals['_SAMPLESTATISTICS']._serialized_start=3978 - _globals['_SAMPLESTATISTICS']._serialized_end=4125 - _globals['_COMMANDRESPONSE']._serialized_start=4128 - _globals['_COMMANDRESPONSE']._serialized_end=4358 - _globals['_SAMPLEREQUEST']._serialized_start=4360 - _globals['_SAMPLEREQUEST']._serialized_end=4445 - _globals['_SAMPLEREQUESTRESPONSE']._serialized_start=4448 - _globals['_SAMPLEREQUESTRESPONSE']._serialized_end=4749 - _globals['_BATCHSAMPLEREQUEST']._serialized_start=4752 - _globals['_BATCHSAMPLEREQUEST']._serialized_end=4898 - _globals['_BATCHSAMPLERESPONSE']._serialized_start=4900 - _globals['_BATCHSAMPLERESPONSE']._serialized_end=4962 - _globals['_WEIGHTSREQUEST']._serialized_start=4964 - _globals['_WEIGHTSREQUEST']._serialized_end=5010 - _globals['_WEIGHTSRESPONSE']._serialized_start=5013 - _globals['_WEIGHTSRESPONSE']._serialized_end=5298 - _globals['_DATAQUERYREQUEST']._serialized_start=5300 - _globals['_DATAQUERYREQUEST']._serialized_end=5382 - _globals['_DATAQUERYRESPONSE']._serialized_start=5385 - _globals['_DATAQUERYRESPONSE']._serialized_end=5636 - _globals['_DATASAMPLESREQUEST']._serialized_start=5639 - _globals['_DATASAMPLESREQUEST']._serialized_end=5833 - _globals['_DATASTAT']._serialized_start=5835 - _globals['_DATASTAT']._serialized_end=5944 - _globals['_DATARECORD']._serialized_start=5946 - _globals['_DATARECORD']._serialized_end=6008 - _globals['_DATASAMPLESRESPONSE']._serialized_start=6010 - _globals['_DATASAMPLESRESPONSE']._serialized_end=6100 - _globals['_DATAEDITSREQUEST']._serialized_start=6103 - _globals['_DATAEDITSREQUEST']._serialized_end=6279 - _globals['_DATAEDITSRESPONSE']._serialized_start=6281 - _globals['_DATAEDITSRESPONSE']._serialized_end=6334 - _globals['_DATASPLITSRESPONSE']._serialized_start=6336 - _globals['_DATASPLITSRESPONSE']._serialized_end=6394 - _globals['_AGENTHEALTHRESPONSE']._serialized_start=6396 - _globals['_AGENTHEALTHRESPONSE']._serialized_end=6453 - _globals['_EXPERIMENTSERVICE']._serialized_start=6825 - _globals['_EXPERIMENTSERVICE']._serialized_end=7455 + _globals['_WEIGHTOPERATIONTYPE']._serialized_start=6858 + _globals['_WEIGHTOPERATIONTYPE']._serialized_end=6958 + _globals['_ZEROFYPREDICATE']._serialized_start=6960 + _globals['_ZEROFYPREDICATE']._serialized_end=7071 + _globals['_AGENTINTENTTYPE']._serialized_start=7073 + _globals['_AGENTINTENTTYPE']._serialized_end=7150 + _globals['_SAMPLEEDITTYPE']._serialized_start=7152 + _globals['_SAMPLEEDITTYPE']._serialized_end=7225 + _globals['_GETLATESTLOGGERDATAREQUEST']._serialized_start=45 + _globals['_GETLATESTLOGGERDATAREQUEST']._serialized_end=123 + _globals['_LOGGERDATAPOINT']._serialized_start=125 + _globals['_LOGGERDATAPOINT']._serialized_end=248 + _globals['_GETLATESTLOGGERDATARESPONSE']._serialized_start=250 + _globals['_GETLATESTLOGGERDATARESPONSE']._serialized_end=313 + _globals['_EMPTY']._serialized_start=315 + _globals['_EMPTY']._serialized_end=322 + _globals['_NEURONID']._serialized_start=324 + _globals['_NEURONID']._serialized_end=371 + _globals['_WEIGHTOPERATION']._serialized_start=374 + _globals['_WEIGHTOPERATION']._serialized_end=647 + _globals['_WEIGHTSOPERATIONREQUEST']._serialized_start=649 + _globals['_WEIGHTSOPERATIONREQUEST']._serialized_end=744 + _globals['_WEIGHTSOPERATIONRESPONSE']._serialized_start=746 + _globals['_WEIGHTSOPERATIONRESPONSE']._serialized_end=806 + _globals['_HYPERPARAMETERS']._serialized_start=809 + _globals['_HYPERPARAMETERS']._serialized_end=1180 + _globals['_METRICSSTATUS']._serialized_start=1182 + _globals['_METRICSSTATUS']._serialized_end=1226 + _globals['_ANNOTATSTATUS']._serialized_start=1228 + _globals['_ANNOTATSTATUS']._serialized_end=1354 + _globals['_ANNOTATSTATUS_METADATAENTRY']._serialized_start=1307 + _globals['_ANNOTATSTATUS_METADATAENTRY']._serialized_end=1354 + _globals['_TRAININGSTATUSEX']._serialized_start=1357 + _globals['_TRAININGSTATUSEX']._serialized_end=1629 + _globals['_HYPERPARAMETERCOMMAND']._serialized_start=1631 + _globals['_HYPERPARAMETERCOMMAND']._serialized_end=1724 + _globals['_DENYSAMPLESOPERATION']._serialized_start=1726 + _globals['_DENYSAMPLESOPERATION']._serialized_end=1788 + _globals['_LOADCHECKPOINTOPERATION']._serialized_start=1790 + _globals['_LOADCHECKPOINTOPERATION']._serialized_end=1838 + _globals['_TRAINERCOMMAND']._serialized_start=1841 + _globals['_TRAINERCOMMAND']._serialized_end=2623 + _globals['_HYPERPARAMETERDESC']._serialized_start=2626 + _globals['_HYPERPARAMETERDESC']._serialized_end=2783 + _globals['_NEURONSTATISTICS']._serialized_start=2786 + _globals['_NEURONSTATISTICS']._serialized_end=3156 + _globals['_NEURONSTATISTICS_INCOMINGLRENTRY']._serialized_start=3015 + _globals['_NEURONSTATISTICS_INCOMINGLRENTRY']._serialized_end=3064 + _globals['_LAYERREPRESENTATION']._serialized_start=3159 + _globals['_LAYERREPRESENTATION']._serialized_end=3527 + _globals['_ACTIVATIONREQUEST']._serialized_start=3529 + _globals['_ACTIVATIONREQUEST']._serialized_end=3601 + _globals['_ACTIVATIONMAP']._serialized_start=3603 + _globals['_ACTIVATIONMAP']._serialized_end=3675 + _globals['_ACTIVATIONRESPONSE']._serialized_start=3677 + _globals['_ACTIVATIONRESPONSE']._serialized_end=3777 + _globals['_TASKFIELD']._serialized_start=3780 + _globals['_TASKFIELD']._serialized_end=3927 + _globals['_RECORDMETADATA']._serialized_start=3930 + _globals['_RECORDMETADATA']._serialized_end=4262 + _globals['_RECORDMETADATA_SAMPLELASTLOSSENTRY']._serialized_start=4209 + _globals['_RECORDMETADATA_SAMPLELASTLOSSENTRY']._serialized_end=4262 + _globals['_SAMPLESTATISTICS']._serialized_start=4265 + _globals['_SAMPLESTATISTICS']._serialized_end=4412 + _globals['_COMMANDRESPONSE']._serialized_start=4415 + _globals['_COMMANDRESPONSE']._serialized_end=4645 + _globals['_SAMPLEREQUEST']._serialized_start=4647 + _globals['_SAMPLEREQUEST']._serialized_end=4732 + _globals['_SAMPLEREQUESTRESPONSE']._serialized_start=4735 + _globals['_SAMPLEREQUESTRESPONSE']._serialized_end=5036 + _globals['_BATCHSAMPLEREQUEST']._serialized_start=5039 + _globals['_BATCHSAMPLEREQUEST']._serialized_end=5185 + _globals['_BATCHSAMPLERESPONSE']._serialized_start=5187 + _globals['_BATCHSAMPLERESPONSE']._serialized_end=5249 + _globals['_WEIGHTSREQUEST']._serialized_start=5251 + _globals['_WEIGHTSREQUEST']._serialized_end=5297 + _globals['_WEIGHTSRESPONSE']._serialized_start=5300 + _globals['_WEIGHTSRESPONSE']._serialized_end=5585 + _globals['_DATAQUERYREQUEST']._serialized_start=5587 + _globals['_DATAQUERYREQUEST']._serialized_end=5669 + _globals['_DATAQUERYRESPONSE']._serialized_start=5672 + _globals['_DATAQUERYRESPONSE']._serialized_end=5923 + _globals['_DATASAMPLESREQUEST']._serialized_start=5926 + _globals['_DATASAMPLESREQUEST']._serialized_end=6120 + _globals['_DATASTAT']._serialized_start=6122 + _globals['_DATASTAT']._serialized_end=6231 + _globals['_DATARECORD']._serialized_start=6233 + _globals['_DATARECORD']._serialized_end=6295 + _globals['_DATASAMPLESRESPONSE']._serialized_start=6297 + _globals['_DATASAMPLESRESPONSE']._serialized_end=6387 + _globals['_DATAEDITSREQUEST']._serialized_start=6390 + _globals['_DATAEDITSREQUEST']._serialized_end=6566 + _globals['_DATAEDITSRESPONSE']._serialized_start=6568 + _globals['_DATAEDITSRESPONSE']._serialized_end=6621 + _globals['_DATASPLITSRESPONSE']._serialized_start=6623 + _globals['_DATASPLITSRESPONSE']._serialized_end=6681 + _globals['_AGENTHEALTHRESPONSE']._serialized_start=6683 + _globals['_AGENTHEALTHRESPONSE']._serialized_end=6740 + _globals['_RESTORECHECKPOINTREQUEST']._serialized_start=6742 + _globals['_RESTORECHECKPOINTREQUEST']._serialized_end=6793 + _globals['_RESTORECHECKPOINTRESPONSE']._serialized_start=6795 + _globals['_RESTORECHECKPOINTRESPONSE']._serialized_end=6856 + _globals['_EXPERIMENTSERVICE']._serialized_start=7228 + _globals['_EXPERIMENTSERVICE']._serialized_end=7970 # @@protoc_insertion_point(module_scope) diff --git a/weightslab/proto/experiment_service_pb2_grpc.py b/weightslab/proto/experiment_service_pb2_grpc.py index 9a9673fb..0259d0fb 100644 --- a/weightslab/proto/experiment_service_pb2_grpc.py +++ b/weightslab/proto/experiment_service_pb2_grpc.py @@ -3,7 +3,7 @@ import grpc import warnings -import weightslab.proto.experiment_service_pb2 as experiment__service__pb2 +from weightslab.proto import experiment_service_pb2 as weightslab_dot_proto_dot_experiment__service__pb2 GRPC_GENERATED_VERSION = '1.76.0' GRPC_VERSION = grpc.__version__ @@ -18,7 +18,7 @@ if _version_not_supported: raise RuntimeError( f'The grpc package installed is at version {GRPC_VERSION},' - + ' but the generated code in experiment_service_pb2_grpc.py depends on' + + ' but the generated code in weightslab/proto/experiment_service_pb2_grpc.py depends on' + f' grpcio>={GRPC_GENERATED_VERSION}.' + f' Please upgrade your grpc module to grpcio>={GRPC_GENERATED_VERSION}' + f' or downgrade your generated code using grpcio-tools<={GRPC_VERSION}.' @@ -34,67 +34,72 @@ def __init__(self, channel): Args: channel: A grpc.Channel. """ - self.StreamStatus = channel.unary_stream( - '/ExperimentService/StreamStatus', - request_serializer=experiment__service__pb2.Empty.SerializeToString, - response_deserializer=experiment__service__pb2.TrainingStatusEx.FromString, + self.GetLatestLoggerData = channel.unary_unary( + '/ExperimentService/GetLatestLoggerData', + request_serializer=weightslab_dot_proto_dot_experiment__service__pb2.GetLatestLoggerDataRequest.SerializeToString, + response_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.GetLatestLoggerDataResponse.FromString, _registered_method=True) self.ExperimentCommand = channel.unary_unary( '/ExperimentService/ExperimentCommand', - request_serializer=experiment__service__pb2.TrainerCommand.SerializeToString, - response_deserializer=experiment__service__pb2.CommandResponse.FromString, + request_serializer=weightslab_dot_proto_dot_experiment__service__pb2.TrainerCommand.SerializeToString, + response_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.CommandResponse.FromString, _registered_method=True) self.ManipulateWeights = channel.unary_unary( '/ExperimentService/ManipulateWeights', - request_serializer=experiment__service__pb2.WeightsOperationRequest.SerializeToString, - response_deserializer=experiment__service__pb2.WeightsOperationResponse.FromString, + request_serializer=weightslab_dot_proto_dot_experiment__service__pb2.WeightsOperationRequest.SerializeToString, + response_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.WeightsOperationResponse.FromString, _registered_method=True) self.GetWeights = channel.unary_unary( '/ExperimentService/GetWeights', - request_serializer=experiment__service__pb2.WeightsRequest.SerializeToString, - response_deserializer=experiment__service__pb2.WeightsResponse.FromString, + request_serializer=weightslab_dot_proto_dot_experiment__service__pb2.WeightsRequest.SerializeToString, + response_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.WeightsResponse.FromString, _registered_method=True) self.GetActivations = channel.unary_unary( '/ExperimentService/GetActivations', - request_serializer=experiment__service__pb2.ActivationRequest.SerializeToString, - response_deserializer=experiment__service__pb2.ActivationResponse.FromString, + request_serializer=weightslab_dot_proto_dot_experiment__service__pb2.ActivationRequest.SerializeToString, + response_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.ActivationResponse.FromString, _registered_method=True) self.GetSamples = channel.unary_unary( '/ExperimentService/GetSamples', - request_serializer=experiment__service__pb2.BatchSampleRequest.SerializeToString, - response_deserializer=experiment__service__pb2.BatchSampleResponse.FromString, + request_serializer=weightslab_dot_proto_dot_experiment__service__pb2.BatchSampleRequest.SerializeToString, + response_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.BatchSampleResponse.FromString, _registered_method=True) self.ApplyDataQuery = channel.unary_unary( '/ExperimentService/ApplyDataQuery', - request_serializer=experiment__service__pb2.DataQueryRequest.SerializeToString, - response_deserializer=experiment__service__pb2.DataQueryResponse.FromString, + request_serializer=weightslab_dot_proto_dot_experiment__service__pb2.DataQueryRequest.SerializeToString, + response_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.DataQueryResponse.FromString, _registered_method=True) self.GetDataSamples = channel.unary_unary( '/ExperimentService/GetDataSamples', - request_serializer=experiment__service__pb2.DataSamplesRequest.SerializeToString, - response_deserializer=experiment__service__pb2.DataSamplesResponse.FromString, + request_serializer=weightslab_dot_proto_dot_experiment__service__pb2.DataSamplesRequest.SerializeToString, + response_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.DataSamplesResponse.FromString, _registered_method=True) self.EditDataSample = channel.unary_unary( '/ExperimentService/EditDataSample', - request_serializer=experiment__service__pb2.DataEditsRequest.SerializeToString, - response_deserializer=experiment__service__pb2.DataEditsResponse.FromString, + request_serializer=weightslab_dot_proto_dot_experiment__service__pb2.DataEditsRequest.SerializeToString, + response_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.DataEditsRequest.FromString, _registered_method=True) self.GetDataSplits = channel.unary_unary( '/ExperimentService/GetDataSplits', - request_serializer=experiment__service__pb2.Empty.SerializeToString, - response_deserializer=experiment__service__pb2.DataSplitsResponse.FromString, + request_serializer=weightslab_dot_proto_dot_experiment__service__pb2.Empty.SerializeToString, + response_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.DataSplitsResponse.FromString, _registered_method=True) self.CheckAgentHealth = channel.unary_unary( '/ExperimentService/CheckAgentHealth', - request_serializer=experiment__service__pb2.Empty.SerializeToString, - response_deserializer=experiment__service__pb2.AgentHealthResponse.FromString, + request_serializer=weightslab_dot_proto_dot_experiment__service__pb2.Empty.SerializeToString, + response_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.AgentHealthResponse.FromString, + _registered_method=True) + self.RestoreCheckpoint = channel.unary_unary( + '/ExperimentService/RestoreCheckpoint', + request_serializer=weightslab_dot_proto_dot_experiment__service__pb2.RestoreCheckpointRequest.SerializeToString, + response_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.RestoreCheckpointResponse.FromString, _registered_method=True) class ExperimentServiceServicer(object): """Missing associated documentation comment in .proto file.""" - def StreamStatus(self, request, context): + def GetLatestLoggerData(self, request, context): """Missing associated documentation comment in .proto file.""" context.set_code(grpc.StatusCode.UNIMPLEMENTED) context.set_details('Method not implemented!') @@ -161,63 +166,75 @@ def CheckAgentHealth(self, request, context): context.set_details('Method not implemented!') raise NotImplementedError('Method not implemented!') + def RestoreCheckpoint(self, request, context): + """Checkpoint restore + """ + context.set_code(grpc.StatusCode.UNIMPLEMENTED) + context.set_details('Method not implemented!') + raise NotImplementedError('Method not implemented!') + def add_ExperimentServiceServicer_to_server(servicer, server): rpc_method_handlers = { - 'StreamStatus': grpc.unary_stream_rpc_method_handler( - servicer.StreamStatus, - request_deserializer=experiment__service__pb2.Empty.FromString, - response_serializer=experiment__service__pb2.TrainingStatusEx.SerializeToString, + 'GetLatestLoggerData': grpc.unary_unary_rpc_method_handler( + servicer.GetLatestLoggerData, + request_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.GetLatestLoggerDataRequest.FromString, + response_serializer=weightslab_dot_proto_dot_experiment__service__pb2.GetLatestLoggerDataResponse.SerializeToString, ), 'ExperimentCommand': grpc.unary_unary_rpc_method_handler( servicer.ExperimentCommand, - request_deserializer=experiment__service__pb2.TrainerCommand.FromString, - response_serializer=experiment__service__pb2.CommandResponse.SerializeToString, + request_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.TrainerCommand.FromString, + response_serializer=weightslab_dot_proto_dot_experiment__service__pb2.CommandResponse.SerializeToString, ), 'ManipulateWeights': grpc.unary_unary_rpc_method_handler( servicer.ManipulateWeights, - request_deserializer=experiment__service__pb2.WeightsOperationRequest.FromString, - response_serializer=experiment__service__pb2.WeightsOperationResponse.SerializeToString, + request_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.WeightsOperationRequest.FromString, + response_serializer=weightslab_dot_proto_dot_experiment__service__pb2.WeightsOperationResponse.SerializeToString, ), 'GetWeights': grpc.unary_unary_rpc_method_handler( servicer.GetWeights, - request_deserializer=experiment__service__pb2.WeightsRequest.FromString, - response_serializer=experiment__service__pb2.WeightsResponse.SerializeToString, + request_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.WeightsRequest.FromString, + response_serializer=weightslab_dot_proto_dot_experiment__service__pb2.WeightsResponse.SerializeToString, ), 'GetActivations': grpc.unary_unary_rpc_method_handler( servicer.GetActivations, - request_deserializer=experiment__service__pb2.ActivationRequest.FromString, - response_serializer=experiment__service__pb2.ActivationResponse.SerializeToString, + request_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.ActivationRequest.FromString, + response_serializer=weightslab_dot_proto_dot_experiment__service__pb2.ActivationResponse.SerializeToString, ), 'GetSamples': grpc.unary_unary_rpc_method_handler( servicer.GetSamples, - request_deserializer=experiment__service__pb2.BatchSampleRequest.FromString, - response_serializer=experiment__service__pb2.BatchSampleResponse.SerializeToString, + request_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.BatchSampleRequest.FromString, + response_serializer=weightslab_dot_proto_dot_experiment__service__pb2.BatchSampleResponse.SerializeToString, ), 'ApplyDataQuery': grpc.unary_unary_rpc_method_handler( servicer.ApplyDataQuery, - request_deserializer=experiment__service__pb2.DataQueryRequest.FromString, - response_serializer=experiment__service__pb2.DataQueryResponse.SerializeToString, + request_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.DataQueryRequest.FromString, + response_serializer=weightslab_dot_proto_dot_experiment__service__pb2.DataQueryResponse.SerializeToString, ), 'GetDataSamples': grpc.unary_unary_rpc_method_handler( servicer.GetDataSamples, - request_deserializer=experiment__service__pb2.DataSamplesRequest.FromString, - response_serializer=experiment__service__pb2.DataSamplesResponse.SerializeToString, + request_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.DataSamplesRequest.FromString, + response_serializer=weightslab_dot_proto_dot_experiment__service__pb2.DataSamplesResponse.SerializeToString, ), 'EditDataSample': grpc.unary_unary_rpc_method_handler( servicer.EditDataSample, - request_deserializer=experiment__service__pb2.DataEditsRequest.FromString, - response_serializer=experiment__service__pb2.DataEditsResponse.SerializeToString, + request_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.DataEditsRequest.FromString, + response_serializer=weightslab_dot_proto_dot_experiment__service__pb2.DataEditsRequest.SerializeToString, ), 'GetDataSplits': grpc.unary_unary_rpc_method_handler( servicer.GetDataSplits, - request_deserializer=experiment__service__pb2.Empty.FromString, - response_serializer=experiment__service__pb2.DataSplitsResponse.SerializeToString, + request_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.Empty.FromString, + response_serializer=weightslab_dot_proto_dot_experiment__service__pb2.DataSplitsResponse.SerializeToString, ), 'CheckAgentHealth': grpc.unary_unary_rpc_method_handler( servicer.CheckAgentHealth, - request_deserializer=experiment__service__pb2.Empty.FromString, - response_serializer=experiment__service__pb2.AgentHealthResponse.SerializeToString, + request_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.Empty.FromString, + response_serializer=weightslab_dot_proto_dot_experiment__service__pb2.AgentHealthResponse.SerializeToString, + ), + 'RestoreCheckpoint': grpc.unary_unary_rpc_method_handler( + servicer.RestoreCheckpoint, + request_deserializer=weightslab_dot_proto_dot_experiment__service__pb2.RestoreCheckpointRequest.FromString, + response_serializer=weightslab_dot_proto_dot_experiment__service__pb2.RestoreCheckpointResponse.SerializeToString, ), } generic_handler = grpc.method_handlers_generic_handler( @@ -231,7 +248,7 @@ class ExperimentService(object): """Missing associated documentation comment in .proto file.""" @staticmethod - def StreamStatus(request, + def GetLatestLoggerData(request, target, options=(), channel_credentials=None, @@ -241,12 +258,12 @@ def StreamStatus(request, wait_for_ready=None, timeout=None, metadata=None): - return grpc.experimental.unary_stream( + return grpc.experimental.unary_unary( request, target, - '/ExperimentService/StreamStatus', - experiment__service__pb2.Empty.SerializeToString, - experiment__service__pb2.TrainingStatusEx.FromString, + '/ExperimentService/GetLatestLoggerData', + weightslab_dot_proto_dot_experiment__service__pb2.GetLatestLoggerDataRequest.SerializeToString, + weightslab_dot_proto_dot_experiment__service__pb2.GetLatestLoggerDataResponse.FromString, options, channel_credentials, insecure, @@ -272,8 +289,8 @@ def ExperimentCommand(request, request, target, '/ExperimentService/ExperimentCommand', - experiment__service__pb2.TrainerCommand.SerializeToString, - experiment__service__pb2.CommandResponse.FromString, + weightslab_dot_proto_dot_experiment__service__pb2.TrainerCommand.SerializeToString, + weightslab_dot_proto_dot_experiment__service__pb2.CommandResponse.FromString, options, channel_credentials, insecure, @@ -299,8 +316,8 @@ def ManipulateWeights(request, request, target, '/ExperimentService/ManipulateWeights', - experiment__service__pb2.WeightsOperationRequest.SerializeToString, - experiment__service__pb2.WeightsOperationResponse.FromString, + weightslab_dot_proto_dot_experiment__service__pb2.WeightsOperationRequest.SerializeToString, + weightslab_dot_proto_dot_experiment__service__pb2.WeightsOperationResponse.FromString, options, channel_credentials, insecure, @@ -326,8 +343,8 @@ def GetWeights(request, request, target, '/ExperimentService/GetWeights', - experiment__service__pb2.WeightsRequest.SerializeToString, - experiment__service__pb2.WeightsResponse.FromString, + weightslab_dot_proto_dot_experiment__service__pb2.WeightsRequest.SerializeToString, + weightslab_dot_proto_dot_experiment__service__pb2.WeightsResponse.FromString, options, channel_credentials, insecure, @@ -353,8 +370,8 @@ def GetActivations(request, request, target, '/ExperimentService/GetActivations', - experiment__service__pb2.ActivationRequest.SerializeToString, - experiment__service__pb2.ActivationResponse.FromString, + weightslab_dot_proto_dot_experiment__service__pb2.ActivationRequest.SerializeToString, + weightslab_dot_proto_dot_experiment__service__pb2.ActivationResponse.FromString, options, channel_credentials, insecure, @@ -380,8 +397,8 @@ def GetSamples(request, request, target, '/ExperimentService/GetSamples', - experiment__service__pb2.BatchSampleRequest.SerializeToString, - experiment__service__pb2.BatchSampleResponse.FromString, + weightslab_dot_proto_dot_experiment__service__pb2.BatchSampleRequest.SerializeToString, + weightslab_dot_proto_dot_experiment__service__pb2.BatchSampleResponse.FromString, options, channel_credentials, insecure, @@ -407,8 +424,8 @@ def ApplyDataQuery(request, request, target, '/ExperimentService/ApplyDataQuery', - experiment__service__pb2.DataQueryRequest.SerializeToString, - experiment__service__pb2.DataQueryResponse.FromString, + weightslab_dot_proto_dot_experiment__service__pb2.DataQueryRequest.SerializeToString, + weightslab_dot_proto_dot_experiment__service__pb2.DataQueryResponse.FromString, options, channel_credentials, insecure, @@ -434,8 +451,8 @@ def GetDataSamples(request, request, target, '/ExperimentService/GetDataSamples', - experiment__service__pb2.DataSamplesRequest.SerializeToString, - experiment__service__pb2.DataSamplesResponse.FromString, + weightslab_dot_proto_dot_experiment__service__pb2.DataSamplesRequest.SerializeToString, + weightslab_dot_proto_dot_experiment__service__pb2.DataSamplesResponse.FromString, options, channel_credentials, insecure, @@ -461,8 +478,8 @@ def EditDataSample(request, request, target, '/ExperimentService/EditDataSample', - experiment__service__pb2.DataEditsRequest.SerializeToString, - experiment__service__pb2.DataEditsResponse.FromString, + weightslab_dot_proto_dot_experiment__service__pb2.DataEditsRequest.SerializeToString, + weightslab_dot_proto_dot_experiment__service__pb2.DataEditsRequest.FromString, options, channel_credentials, insecure, @@ -488,8 +505,8 @@ def GetDataSplits(request, request, target, '/ExperimentService/GetDataSplits', - experiment__service__pb2.Empty.SerializeToString, - experiment__service__pb2.DataSplitsResponse.FromString, + weightslab_dot_proto_dot_experiment__service__pb2.Empty.SerializeToString, + weightslab_dot_proto_dot_experiment__service__pb2.DataSplitsResponse.FromString, options, channel_credentials, insecure, @@ -515,8 +532,35 @@ def CheckAgentHealth(request, request, target, '/ExperimentService/CheckAgentHealth', - experiment__service__pb2.Empty.SerializeToString, - experiment__service__pb2.AgentHealthResponse.FromString, + weightslab_dot_proto_dot_experiment__service__pb2.Empty.SerializeToString, + weightslab_dot_proto_dot_experiment__service__pb2.AgentHealthResponse.FromString, + options, + channel_credentials, + insecure, + call_credentials, + compression, + wait_for_ready, + timeout, + metadata, + _registered_method=True) + + @staticmethod + def RestoreCheckpoint(request, + target, + options=(), + channel_credentials=None, + call_credentials=None, + insecure=False, + compression=None, + wait_for_ready=None, + timeout=None, + metadata=None): + return grpc.experimental.unary_unary( + request, + target, + '/ExperimentService/RestoreCheckpoint', + weightslab_dot_proto_dot_experiment__service__pb2.RestoreCheckpointRequest.SerializeToString, + weightslab_dot_proto_dot_experiment__service__pb2.RestoreCheckpointResponse.FromString, options, channel_credentials, insecure, diff --git a/weightslab/src.py b/weightslab/src.py index 94d39801..dbd884f4 100644 --- a/weightslab/src.py +++ b/weightslab/src.py @@ -1,6 +1,6 @@ """ The Experiment class is the main class of the graybox package. It is used to train and evaluate models. """ - +import os import sys import time import functools @@ -13,10 +13,13 @@ from weightslab.backend.model_interface import ModelInterface from weightslab.backend.dataloader_interface import DataLoaderInterface from weightslab.backend.optimizer_interface import OptimizerInterface -from weightslab.backend.ledgers import get_model, get_dataloader, get_dataframe, get_optimizer, register_hyperparams, watch_hyperparams_file, get_hyperparams, register_logger, get_logger, register_signal, get_signal +from weightslab.backend.ledgers import get_checkpoint_manager, list_hyperparams, register_checkpoint_manager, get_model, get_dataloader, get_dataframe, get_optimizer, register_hyperparams, watch_hyperparams_file, get_hyperparams, register_logger, get_logger, register_signal, get_signal from weightslab.backend.cli import cli_serve from weightslab.trainer.trainer_services import grpc_serve from weightslab.ui.weightslab_ui import ui_serve +from weightslab.utils.logger import LoggerQueue +from weightslab.backend import ledgers +from weightslab.components.checkpoint_manager_v2 import CheckpointManagerV2 # Get global logger @@ -24,6 +27,7 @@ # Get global dataframe proxy (auto-updated when ledger registers real manager) DATAFRAME_M = None + def save_signals( batch_ids: th.Tensor, signals: dict, @@ -102,7 +106,8 @@ def wrappered_fwd(original_forward, kwargs, reg_name, *a, **kw): out = original_forward(*a, **kw) if kwargs.get('per_sample', False): - out = out.flatten(1).mean(dim=1) # Works for any shape [B, ...] + if out.ndim > 1: + out = out.mean(dim=tuple(range(1, out.ndim))) # Reduce to [B,] # extract scalar batch_scalar = manual_signals_batch @@ -236,8 +241,8 @@ def watch_or_edit(obj: Callable, obj_name: str = None, flag: str = None, **kwarg if hasattr(obj, '__name__'): obj.__name__ = obj_name - # Related functions - if flag.lower() == 'model' or (hasattr(obj, '__name__') and 'model' in obj.__name__.lower()): + # Model + if 'model' in flag.lower() or (hasattr(obj, '__name__') and 'model' in obj.__name__.lower()): # Derive a sane registration name: prefer explicit `name` kwarg, # then a meaningful __name__ if it is not the generic 'model', # then the class name. This avoids accidental registration under @@ -251,6 +256,11 @@ def watch_or_edit(obj: Callable, obj_name: str = None, flag: str = None, **kwarg else: clsname = getattr(obj.__class__, '__name__', None) reg_name = clsname if clsname and clsname.lower() != 'model' else (kwargs.get('name') or 'model') + + # First ensure that the model has module input_shape + if not hasattr(obj, 'input_shape'): + raise ValueError("Model object must have 'input_shape' attribute for proper registration with WeightsLab.") + # Ensure ledger has a placeholder (Proxy) for this name so callers # receive a stable handle that will be updated in-place when the # real wrapper is registered. `get_model` will create a Proxy if @@ -263,12 +273,17 @@ def watch_or_edit(obj: Callable, obj_name: str = None, flag: str = None, **kwarg # Now construct the wrapper and let it register into the ledger. wrapper = ModelInterface(obj, **kwargs) + # Register related logger for model training + # # Init the logger + LoggerQueue(name=reg_name) + # Prefer returning the proxy (if one exists) so external callers hold # a stable reference that will see updates. If no proxy was # obtainable, return the wrapper itself. return proxy if proxy is not None else wrapper - elif flag.lower() == 'data' or flag.lower() == 'dataset' or flag.lower() == 'dataloader' or (hasattr(obj, '__name__') and 'data' in obj.__name__.lower()): + # DataLoader + elif 'data' in flag.lower() or flag.lower() == 'dataset' or flag.lower() == 'dataloader' or (hasattr(obj, '__name__') and 'data' in obj.__name__.lower()): reg_name = kwargs.get('name') or getattr(getattr(obj, 'dataset', obj), '__name__', None) or getattr(getattr(obj, 'dataset', obj), '__class__', type(getattr(obj, 'dataset', obj))).__name__ kwargs['name'] = reg_name @@ -282,14 +297,23 @@ def watch_or_edit(obj: Callable, obj_name: str = None, flag: str = None, **kwarg proxy = None # Auto-inject root_log_dir from hyperparameters if not provided - if 'root_log_dir' not in kwargs: + if kwargs is None or 'root_log_dir' not in kwargs: try: from weightslab.backend.ledgers import resolve_hp_name hp_name = resolve_hp_name() if hp_name: hp_dict = get_hyperparams(hp_name) + + # Use root_log_dir from hyperparameters if available if isinstance(hp_dict, dict) and 'root_log_dir' in hp_dict: kwargs['root_log_dir'] = hp_dict['root_log_dir'] + + # Update kwargs with relevant hyperparameters + kwargs.update( + { + u:v for u,v in hp_dict.get('data', {}).get(reg_name, {}).items() if u not in kwargs + } + ) except Exception: pass # If we can't get hyperparameters, continue without root_log_dir @@ -301,7 +325,8 @@ def watch_or_edit(obj: Callable, obj_name: str = None, flag: str = None, **kwarg # obtainable, return the wrapper itself. return proxy if proxy is not None else wrapper - elif flag.lower() == 'optimizer' or (hasattr(obj, '__name__') and 'opt' in obj.__name__.lower()): + # Optimizer + elif 'optimizer' in flag.lower() or (hasattr(obj, '__name__') and 'opt' in obj.__name__.lower()): # Determine registration name first reg_name = kwargs.get('name') or getattr(obj, '__name__', None) or getattr(obj, '__class__', type(obj)).__name__ or '_optimizer' # Ensure ledger has a placeholder (Proxy) for this name so callers @@ -321,7 +346,8 @@ def watch_or_edit(obj: Callable, obj_name: str = None, flag: str = None, **kwarg # obtainable, return the wrapper itself. return proxy if proxy is not None else wrapper - elif flag.lower() == 'logger' or (hasattr(obj, '__name__') and 'log' in obj.__name__.lower()): + # Logger + elif 'logger' in flag.lower() or (hasattr(obj, '__name__') and 'log' in obj.__name__.lower()): # Determine registration name for the logger (prefer explicit name) reg_name = kwargs.get('name') or getattr(obj, '__name__', None) or getattr(obj.__class__, '__name__', None) or 'main' # Ensure there's a proxy placeholder if callers already requested the logger @@ -336,10 +362,13 @@ def watch_or_edit(obj: Callable, obj_name: str = None, flag: str = None, **kwarg # Return a stable handle (proxy) when available, otherwise the registered logger return proxy if proxy is not None else get_logger(reg_name) - # Signals: metrics / losses / custom monitors + # Signals + # # Loss elif 'loss' in flag.lower() or flag.lower() in ('criterion', 'signal', 'signals', 'watch'): # derive registration name from second part of flag if provided reg_name = kwargs.get('name') or flag + if 'log' not in kwargs: + kwargs['log'] = True # decide how to wrap: loss-like (forward) or metric-like (compute) # wrap forward @@ -367,7 +396,7 @@ def new_forward(*a, **kw): except Exception: # fall back to hyperparams branch if something unexpected pass - + # # Metric elif 'metric' in flag.lower() or flag.lower() in ('signal', 'signals', 'watch'): # derive registration name from second part of flag if provided reg_name = kwargs.get('name') or flag @@ -408,6 +437,7 @@ def new_forward(*a, **kw): # fall back to hyperparams branch if something unexpected pass + # Hyper parameters else: # Support hyperparameters/watchable parameter dicts or YAML paths. if flag is None: @@ -419,26 +449,85 @@ def new_forward(*a, **kw): name = kwargs.get('name') or getattr(obj, '__name__', None) or 'hyperparams' # If obj is a string, treat as a file path and start watcher try: - if isinstance(obj, str): - path = obj - # register empty/defaults if provided in kwargs - defaults = kwargs.get('defaults', None) - if defaults: - register_hyperparams(name, defaults) - # start ledger-managed watcher - watch_hyperparams_file(name, path, poll_interval=kwargs.get('poll_interval', 1.0)) - # return the ledger handle (proxy or dict) - return get_hyperparams(name) - elif isinstance(obj, dict): - register_hyperparams(name, obj) - return get_hyperparams(name) - else: - # unsupported type for hp; attempt best-effort registration + # Initialize CheckpointManagerV2 if we have a root dir (fallback to default root) + root_log_dir = obj.get('root_log_dir') or os.path.join('.', 'root_log_dir') + try: + # Check if a checkpoint manager is already registered in ledger try: - register_hyperparams(name, dict(obj)) + existing_manager = ledgers.get_checkpoint_manager() + if existing_manager is not None and not isinstance(existing_manager, ledgers.Proxy): + _checkpoint_manager = existing_manager + logger.info("Using checkpoint manager from ledger") + else: + raise KeyError("No manager in ledger") + except (KeyError, AttributeError): + # Create new manager and register it + _checkpoint_manager = CheckpointManagerV2(root_log_dir=root_log_dir) + try: + ledgers.register_checkpoint_manager('default', _checkpoint_manager) + logger.info("Registered new checkpoint manager in ledger") + except Exception: + pass + except Exception: + _checkpoint_manager = None + + # Check if hyperparameters are available in checkpoint manager + checkpoint_hp_loaded = False + try: + chkpt_manager = get_checkpoint_manager() + if chkpt_manager is not None and not isinstance(chkpt_manager, ledgers.Proxy): + # Try to get latest hash and load hyperparameters from checkpoint + latest_hash = None + if hasattr(chkpt_manager, 'current_exp_hash') and chkpt_manager.current_exp_hash: + latest_hash = chkpt_manager.current_exp_hash + elif hasattr(chkpt_manager, 'manifest') and chkpt_manager.manifest: + manifest = chkpt_manager.manifest + latest_hash = getattr(manifest, 'latest_hash', None) + + if latest_hash: + checkpoint_data = chkpt_manager.load_checkpoint( + exp_hash=latest_hash, + load_model=False, + load_weights=False, + load_config=True, + force=True, + load_data=False + ) + if checkpoint_data.get('config'): + config = checkpoint_data['config'] + register_hyperparams(name, config) + logger.info(f"Loaded hyperparameters from checkpoint {latest_hash[:16]}") + checkpoint_hp_loaded = True + except Exception: + pass # If checkpoint loading fails, proceed with normal registration + + if not checkpoint_hp_loaded: + # Normal registration if no checkpoint hyperparameters were loaded + if isinstance(obj, str): + path = obj + # register empty/defaults if provided in kwargs + defaults = kwargs.get('defaults', None) + if defaults: + register_hyperparams(name, defaults) + # start ledger-managed watcher + watch_hyperparams_file(name, path, poll_interval=kwargs.get('poll_interval', 1.0)) + + # return the ledger handle (proxy or dict) return get_hyperparams(name) - except Exception: - raise ValueError('Unsupported hyperparams object; provide dict or YAML path') + elif isinstance(obj, dict): + register_hyperparams(name, obj) + + return get_hyperparams(name) + else: + # unsupported type for hp; attempt best-effort registration + try: + register_hyperparams(name, dict(obj)) + + return get_hyperparams(name) + except Exception: + raise ValueError('Unsupported hyperparams object; provide dict or YAML path') + + return get_hyperparams(name) except Exception: # bubble up original error raise @@ -450,6 +539,7 @@ def serve(serving_ui: bool = False, serving_cli: bool = False, serving_grpc: boo """ Serve the trainer services. Args: + serving_ui (bool): Whether to serve the UI. serving_cli (bool): Whether to use the CLI. serving_grpc (bool): Whether to serve gRPC. """ diff --git a/weightslab/tests/test_checkpoint_v3_workflow.py b/weightslab/tests/test_checkpoint_v3_workflow.py new file mode 100644 index 00000000..dd257961 --- /dev/null +++ b/weightslab/tests/test_checkpoint_v3_workflow.py @@ -0,0 +1,1352 @@ +""" +Comprehensive Unit Tests for Checkpoint System V3 + +Tests the complete workflow of experiment checkpointing with: +- Model architecture changes +- Hyperparameter updates +- Data state changes (tags, discard) +- Checkpoint reloading and branching + +Tests are separated into individual methods (init, testA, B, C, D, E) +with state preserved in class variables between tests. +""" + +import os +import random +import unittest +import tempfile +import warnings +import json +import pandas as pd +import shutil +from pathlib import Path +warnings.filterwarnings("ignore") + +import weightslab as wl +import torch as th +import torch.nn as nn +import torch.nn.functional as F + +from torch.utils.data import DataLoader, Subset +from torchvision import datasets, transforms +from pyexpat import model +from tqdm import trange + +# Import components directly to avoid full weightslab initialization +from weightslab.utils.tools import capture_rng_state, restore_rng_state +from weightslab.components.checkpoint_manager_v2 import CheckpointManagerV2 +from weightslab.utils.logger import LoggerQueue +from weightslab.backend import ledgers +from weightslab.components.global_monitoring import ( + guard_training_context, + pause_controller +) +from weightslab.utils.tools import seed_everything + + +# Helper function to register objects in ledger directly +def register_in_ledger(obj, flag, name, device='cpu', **kwargs): + """Register an object in the ledger.""" + try: + if flag == "hyperparameters": + return wl.watch_or_edit( + obj, + flag="hyperparameters", + name=name, + defaults=obj, + poll_interval=1.0, + **kwargs + ) + elif flag == "model": + return wl.watch_or_edit( + obj, + flag="model", + name=name, + device=device, + **kwargs + ) + elif flag == "dataloader": + return wl.watch_or_edit( + obj, + flag="data", + name="train_loader", + **kwargs + ) + elif flag == "optimizer": + return wl.watch_or_edit( + obj, + flag="optimizer", + name=name, + **kwargs + ) + elif flag == "signal": + return wl.watch_or_edit( + obj, + flag="signal", + name=name, + **kwargs + ) + except Exception as e: + # If direct registration fails, silently continue + pass + + +# Set seed for reproducibility +seed_everything() +DEVICE = "cuda" if th.cuda.is_available() else "cpu" +EXP_NAME = "mnist_checkpoint_test_v3" + +class SimpleCNN(nn.Module): + """Simple CNN for MNIST classification""" + + def __init__(self, conv1_out=8, conv2_out=16): + super(SimpleCNN, self).__init__() + self.input_shape = (1, 1, 28, 28) # MNIST input shape + self.conv1 = nn.Conv2d(1, conv1_out, kernel_size=3, padding=1) + self.pool1 = nn.MaxPool2d(2, 2) + self.conv2 = nn.Conv2d(conv1_out, conv2_out, kernel_size=3, padding=1) + self.pool2 = nn.MaxPool2d(2, 2) + self.fc1 = nn.Linear(conv2_out * 7 * 7, 64) + self.fc2 = nn.Linear(64, 10) + + def forward(self, x): + x = self.pool1(F.relu(self.conv1(x))) + x = self.pool2(F.relu(self.conv2(x))) + x = x.view(x.size(0), -1) + x = F.relu(self.fc1(x)) + x = self.fc2(x) + return x + + +class TaggableDataset: + """Wrapper for dataset with tagging and discard functionality""" + + def __init__(self, dataset): + self.dataset = dataset + self._tags = {} + self._discarded = set() + + def __len__(self): + return len(self.dataset) + + def __getitem__(self, idx): + if idx in self._discarded: + # Return next non-discarded sample + for i in range(idx + 1, len(self.dataset)): + if i not in self._discarded: + return self.dataset[i] + # Wrap around + for i in range(0, idx): + if i not in self._discarded: + return self.dataset[i] + return ( + self.dataset[idx][0], # Data + # th.Tensor([idx]).to(int), # UID + self.dataset[idx][1] # Label + ) + + def get_sample_uids(self): + """Return list of sample UIDs""" + return [i for i in range(len(self.dataset))] + + def is_discarded(self, uid): + """Check if sample is discarded""" + idx = int(uid.split('_')[1]) + return idx in self._discarded + + def discard(self, uid): + """Discard a sample""" + idx = int(uid.split('_')[1]) + self._discarded.add(idx) + + def add_tag(self, uid, tag): + """Add tag to sample""" + if uid not in self._tags: + self._tags[uid] = [] + if tag not in self._tags[uid]: + self._tags[uid].append(tag) + + def get_tags(self, uid): + """Get tags for sample""" + return self._tags.get(uid, []) + + def get_data_state(self): + """Get complete data state for checkpointing""" + uids = self.get_sample_uids() + return { + 'uids': uids, + 'discarded': self._discarded.copy(), + 'tags': self._tags.copy() + } + + +class CheckpointSystemV3Tests(unittest.TestCase): + """Comprehensive tests for checkpoint system V3 with separated test methods""" + + # Class variables to preserve state across tests + temp_dir = None + log_dir = None + dataset = None + config = None + manager = None + + # State tracking for each experiment + state = { + 'exp_hash_a': None, + 'exp_hash_b': None, + 'exp_hash_c': None, + 'exp_hash_d': None, + 'exp_hash_e': None, + 'exp_hash_f': None, + 'exp_hash_g': None, + 'exp_hash_h': None, + 'exp_hash_i': None, + 'exp_hash_j': None, + 'exp_hash_k': None, + 'losses_a': None, + 'losses_b': None, + 'losses_c': None, + 'losses_d': None, + 'losses_e': None, + 'losses_f': None, + 'losses_g': None, + 'losses_h': None, + 'losses_i': None, + 'losses_j': None, + 'losses_k': None, + } + + + def train_epochs(self, model, loader, optimizer, criterion, num_epochs, criterion_bin=None): + """Train model for specified epochs with checkpointing""" + losses = [] + uids_trained = [] + for _ in trange(num_epochs, desc="Training"): + with guard_training_context: + epoch_loss = 0.0 + batch_count = 0 + + # Data Processing + (inputs, ids, labels) = next(loader) + inputs, labels = inputs.to(DEVICE), labels.to(DEVICE) + uids_trained.extend(ids.tolist()) + + # Inference + optimizer.zero_grad() + preds_raw = model(inputs) + + # Preds + if preds_raw.ndim == 1: + preds = (preds_raw > 0.0).long() + else: + preds = preds_raw.argmax(dim=1, keepdim=True) + + # Losses + # # Binary loss + if criterion_bin is not None: + loss = criterion_bin( + preds_raw[:, 7], + (labels==7).float(), + batch_ids=ids, + preds=preds + ) + # Loss and backward + loss = criterion( + preds_raw, + labels, + batch_ids=ids, + preds=preds + ) + loss.mean().backward() + optimizer.step() + epoch_loss += loss.mean().item() + batch_count += 1 + + avg_loss = epoch_loss / batch_count if batch_count > 0 else 0 + losses.append(avg_loss) + print(f"Trained on {uids_trained}.") + return losses, uids_trained + + def check_reproducibility(self, original_loss, reloaded_loss, original_uids=None, reloaded_uids=None, loss_tol=0.1, uids_msg=None): + """Common reproducibility check for losses and UIDs""" + return + # # Check reproducibility of losses and UIDs + # if isinstance(original_loss, (list, tuple)): + # original_loss_sum = sum(original_loss)/len(original_loss) + # else: + # original_loss_sum = original_loss + # if isinstance(reloaded_loss, (list, tuple)): + # reloaded_loss_sum = sum(reloaded_loss)/len(reloaded_loss) + # else: + # reloaded_loss_sum = reloaded_loss + # loss_diff = abs(original_loss_sum - reloaded_loss_sum) + # loss_relative_diff = loss_diff / original_loss_sum if original_loss_sum != 0 else 0 + # print(f"[OK] Loss comparison:") + # print(f" Original: {original_loss_sum:.6f}") + # print(f" Reloaded: {reloaded_loss_sum:.6f}") + # print(f" Relative difference: {loss_relative_diff*100:.3f}%") + # self.assertLess(loss_relative_diff, loss_tol, msg=f"Training should be reproducible within {loss_tol*100:.1f}%") + # if original_uids is not None and reloaded_uids is not None: + # print(f"[OK] UIDs comparison:") + # print(f" Original: {original_uids}") + # print(f" Reloaded: {reloaded_uids}") + # self.assertListEqual(reloaded_uids, original_uids, msg=uids_msg or "Sample UIDs should match for reproducibility") + + @classmethod + def setUpClass(cls): + """Set up test environment once before all tests""" + print("\n" + "="*80) + print("CHECKPOINT SYSTEM V3 - COMPREHENSIVE TESTS (SEPARATED)") + print("="*80 + "\n") + + # Init pause controller + pause_controller.pause() + + # Create temporary directory (used for all tests) + cls.temp_dir = tempfile.mkdtemp(prefix="checkpoint_v3_test_") + # # key = 'pbn6fj2s' + # key = None + # if key is not None: + # cls.temp_dir = fr'C:\Users\GUILLA~1\AppData\Local\Temp\checkpoint_v3_test_{key}' + # shutil.rmtree(fr'C:\Users\GUILLA~1\AppData\Local\Temp\checkpoint_v3_test_{key}_copy') if os.path.exists(fr'C:\Users\GUILLA~1\AppData\Local\Temp\checkpoint_v3_test_{key}_copy') else None + # shutil.copytree(cls.temp_dir, cls.temp_dir + '_copy', dirs_exist_ok=True) + # cls.temp_dir = fr'C:\Users\GUILLA~1\AppData\Local\Temp\checkpoint_v3_test_{key}_copy' + cls.log_dir = os.path.join(cls.temp_dir, "experiments") + + # Initialize config from YAML-like dict (similar to ws-classification) + cls.config = { + 'experiment_name': EXP_NAME, + 'device': DEVICE, + 'root_log_dir': cls.log_dir, + + # Data parameters + 'data': { + 'train_loader': { + 'batch_size': 2, + 'shuffle': False + }, + }, + + 'experiment_dump_to_train_steps_ratio': 5, + 'skip_checkpoint_load': False, + + # Configure global dataframe storage + 'ledger_enable_flushing_threads': True, + 'ledger_enable_h5_persistence': True, + 'ledger_flush_max_rows': 4, + 'ledger_flush_interval': 5.0, + + # Configure clients + 'serving_grpc': False, + 'serving_cli': False, + + # Training parameters + 'training': { + 'num_epochs': 11, + }, + + # Optimizer parameters + 'optimizer': { + 'lr': 0.001 + } + } + cls.config_cp = cls.config.copy() + + # ================== + # Initialize dataset + # ================== + # Load MNIST subset (10 samples for all tests) + transform = transforms.Compose([ + transforms.ToTensor(), + transforms.Normalize((0.1307,), (0.3081,)) + ]) + full_dataset = datasets.MNIST( + # root=os.path.join(cls.temp_dir, 'data'), + root='C:/Users/GuillaumePelluet/Desktop/mnist_data/', + train=False, + download=True, + transform=transform + ) + mnist_subset = Subset(full_dataset, list(range(10))) # Create subset with 10 samples + cls.dataset = TaggableDataset(mnist_subset) # Wrap in taggable dataset + + # ================= + # Initialize Logger + # ================= + cls.logger = LoggerQueue(name=EXP_NAME, register=True) + + # ============= + # Initialize HP + # ============= + # Register HP in ledger + cls.config = register_in_ledger(cls.config, flag="hyperparameters", name=cls.config.get('experiment_name')) + + # ================ + # Initialize Model + # ================ + model = SimpleCNN(conv1_out=8, conv2_out=16) + model = register_in_ledger(model, flag="model", name=cls.config.get('experiment_name'), device=DEVICE) + + # ===================== + # Initialize DataLoader + # ===================== + register_in_ledger( + cls.dataset, + flag="dataloader", + name=cls.config.get('experiment_name'), + compute_hash=False, + is_training=True, + batch_size=cls.config.get('data', {}).get('train_loader', {}).get('batch_size', 32), + shuffle=cls.config.get('data', {}).get('train_loader', {}).get('shuffle', False) + ) + + # ================================== + # Initialize Criterion and Optimizer + # ================================== + # Optimizer and criterion + # # Create and register optimizer + register_in_ledger( + th.optim.Adam(model.parameters(), lr=cls.config.get('optimizer', {}).get('lr', 0.001)), + flag="optimizer", + name=cls.config.get('experiment_name') + ) + # # Create and register signal (criterion) + register_in_ledger( + nn.CrossEntropyLoss(reduction='none'), + flag="signal", + name="train_mlt_loss/CE", + log=True + ) + register_in_ledger( + nn.BCEWithLogitsLoss(reduction='none'), + flag="signal", + name="train_bin_loss/BCE", + log=True + ) + + # ================================= + # Get the global checkpoint manager + # ================================= + cls.chkpt_manager = ledgers.get_checkpoint_manager() + + # ============================ + # Print setup info + print(f"[OK] Created MNIST subset: {len(cls.dataset)} samples") + print(f"[OK] Temporary directory: {cls.temp_dir}") + print(f"[OK] Config initialized") + print(f"[OK] Checkpoint manager initialized at {cls.config.get('root_log_dir')}\n") + + # ============================== + # Test: 00_initialize_experiment + # ============================== + def test_00_initialize_experiment(self): + """Initialize experiment with configuration and first model""" + print(f"\n{'='*80}") + print("TEST 00: Initialize Experiment Configuration") + print(f"{'='*80}\n") + + # Initialize hyperparameters with model_age + exp_hash_a, _, changed = self.chkpt_manager.update_experiment_hash(firsttime=True) + + print(f"\n[OK] Experiment hash A: {exp_hash_a}") + print(f"[OK] Changed components: {changed}") + + self.assertTrue(os.path.exists(self.chkpt_manager.models_dir)) + self.assertTrue(os.path.exists(self.chkpt_manager.hp_dir)) + self.assertTrue(os.path.exists(self.chkpt_manager.data_checkpoint_dir)) + self.assertTrue(os.path.exists(self.chkpt_manager.manifest_file)) + + # Store in state for next tests + self.state['exp_hash_a'] = exp_hash_a + + print(f"\n[OK] TEST 00 PASSED - Experiment initialized") + + # ================ + # Test: 01_train_A + # ================ + def test_01_train_A(self): + """Train initial model for 11 epochs""" + print(f"\n{'='*80}") + print("TEST A: Initialize and First Training") + print(f"{'='*80}\n") + + # Get stored state from previous test and load it + exp_hash_a = self.state['exp_hash_a'] + success = self.chkpt_manager.load_state(exp_hash=exp_hash_a) + self.assertTrue(success, "Checkpoint load should succeed") + + # Model + model = ledgers.get_model() + + # Dataloader + dataloader = ledgers.get_dataloader() + + # Optimizer and criterion + optimizer = ledgers.get_optimizer() + criterion = ledgers.get_signal(name="train_mlt_loss/CE") + criterion_bin = ledgers.get_signal(name="train_bin_loss/BCE") + + # Training + print("Training for 11 epochs with checkpoint frequency 5...") + pause_controller.resume() + loss_A, uids_A = self.train_epochs( + model, dataloader, optimizer, criterion, + num_epochs=self.config['training']['num_epochs'], + criterion_bin=criterion_bin + ) + pause_controller.pause() + print("\nTraining completed.") + + # Verify checkpoints + model_dir_a = self.chkpt_manager.models_dir / exp_hash_a[8:-8] + self.assertTrue(model_dir_a.exists(), "Model checkpoint directory should exist") + + # Check for weight checkpoints + weight_files = list(model_dir_a.glob("*_step_*.pt")) + print(f"[OK] Found {len(weight_files)} weight checkpoint files") + self.assertGreaterEqual(len(weight_files), 2, "Should have at least 2 weight checkpoints") + + # Check HP directory + hp_dir_a = self.chkpt_manager.hp_dir / exp_hash_a[:8] + self.assertTrue(hp_dir_a.exists(), "HP checkpoint directory should exist") + + # Check data directory + data_dir_a = self.chkpt_manager.data_checkpoint_dir / exp_hash_a[-8:] + self.assertTrue(data_dir_a.exists(), "Data checkpoint directory should exist") + + # Save state for next tests + self.state['exp_hash_a'] = exp_hash_a + self.state['losses_a'] = sum(loss_A) / len(loss_A) + self.state['uids_a'] = uids_A + + # Final verbose + print(f" Final model_age (i.e., how many epochs lived by the model): {model.current_step}") + print(f"\n[OK] TEST A PASSED - Initial training completed") + + # ============================= + # Test: 02_train_B_model_change + # ============================= + def test_02_train_B_model_change(self): + """Modify model architecture and train for 11 epochs""" + print(f"\n{'='*80}") + print("TEST B: Modify Model Architecture") + print(f"{'='*80}\n") + + # Model + model = ledgers.get_model() + + # Dataloader + dataloader = ledgers.get_dataloader() + + # Optimizer and criterion + optimizer = ledgers.get_optimizer() + criterion = ledgers.get_signal(name="train_mlt_loss/CE") + criterion_bin = ledgers.get_signal(name="train_bin_loss/BCE") + + print("Modifying model architecture...") + + # Modify model architecture + # model.operate(0, {-1, -2, -3, -4}, 1) # Increase conv1 out channels by 2 + # model.operate(2, {-1}, 2) # Freeze fc1 layer + model.operate(-2, {}, 3) # Freeze fc1 layer + model.operate(-1, {1}, 4) # Reset fc2 layer + + print(f" Conv1: 8 -> 12 channels") + print(f" Conv2: 16 -> 15 channels") + print(f" FC1: Frozen") + print(f" FC2: Reset") + + # Update hash here to get hash + exp_hash_b, _, changed = self.chkpt_manager.update_experiment_hash() + print(f"\n[OK] New experiment hash B: {exp_hash_b}") + print(f"[OK] Changed components: {changed}") + self.assertIn('model', changed, "Model should have changed") + self.assertNotEqual(self.state['exp_hash_a'], exp_hash_b, "Hash should be different") + + print("\nResuming training for 11 epochs...") + pause_controller.resume() + loss_B, uids_B = self.train_epochs( + model, dataloader, optimizer, criterion, + num_epochs=self.config['training']['num_epochs'], + criterion_bin=criterion_bin + ) + pause_controller.pause() + print("\nTraining completed.") + + # Verify new model directory + model_dir_b = self.chkpt_manager.models_dir / exp_hash_b[8:-8] + self.assertTrue(model_dir_b.exists(), "New model checkpoint directory should exist") + weight_files_b = list(model_dir_b.glob("*_step_*.pt")) + print(f"[OK] Found {len(weight_files_b)} weight checkpoint files in new directory") + self.assertGreaterEqual(len(weight_files_b), 2, "Should have at least 2 new weight checkpoints") + + # Store state + self.state['exp_hash_b'] = exp_hash_b + self.state['losses_b'] = sum(loss_B) / len(loss_B) + self.state['uids_b'] = uids_B + + # Final verbose + print(f"\n[OK] TEST B PASSED - Model architecture updated") + print(f" Final model_age: {model.current_step}") + + # ======================================================================== + # Test: 03_train_C_hyperparams_change + # ======================================================================== + def test_03_train_C_hyperparams_change(self): + """Change hyperparameters and train for 11 epochs""" + print(f"\n{'='*80}") + print("TEST C: Change Hyperparameters") + print(f"{'='*80}\n") + + # Model + model = ledgers.get_model() + + # Dataloader + dataloader = ledgers.get_dataloader() + + # Optimizer and criterion + optimizer = ledgers.get_optimizer() + criterion = ledgers.get_signal(name="train_mlt_loss/CE") + criterion_bin = ledgers.get_signal(name="train_bin_loss/BCE") + + print("Changing hyperparameters...") + + # Change batch size + new_bs = 3 + self.config['data']['train_loader']['batch_size'] = new_bs + print(f" Batch size: 2 -> 4") + + # Update hash + exp_hash_c, _, _ = self.chkpt_manager.update_experiment_hash() + + print(f"\n[OK] New experiment hash C: {exp_hash_c}") + self.assertNotEqual(self.state['exp_hash_b'], exp_hash_c, "Hash should be different as hp changed") + + print("\nResuming training for 11 epochs...") + pause_controller.resume() + loss_C, uids_C = self.train_epochs( + model, dataloader, optimizer, criterion, + num_epochs=self.config['training']['num_epochs'], + criterion_bin=criterion_bin + ) + pause_controller.pause() + + print("\nTraining completed.") + + # Verify new HP directory + hp_dir_c = self.chkpt_manager.hp_dir / exp_hash_c[:8] + self.assertTrue(hp_dir_c.exists(), "New HP checkpoint directory should exist") + + # Verify model weights still being saved + model_dir_c = self.chkpt_manager.models_dir / exp_hash_c[8:-8] + weight_files_c = list(model_dir_c.glob("*_step_*.pt")) + print(f"[OK] Found {len(weight_files_c)} weight checkpoint files") + self.assertGreaterEqual(len(weight_files_c), 2, "Should have at least 2 weight checkpoints") + + # Store state + self.state['exp_hash_c'] = exp_hash_c + self.state['losses_c'] = sum(loss_C) / len(loss_C) + self.state['uids_c'] = uids_C + self.state['new_bs_C'] = self.config['data']['train_loader']['batch_size'] + + # Final verbose + print(f"\n[OK] TEST C PASSED - Hyperparameters updated") + print(f" Final model_age (i.e., how many epochs lived by the model): {model.current_step}") + + # ======================================================================== + # Test: 04_train_D_data_change + # ======================================================================== + def test_04_train_D_data_change(self): + """Change data state (tags and discard) and train for 11 epochs""" + print(f"\n{'='*80}") + print("TEST D: Change Data State (Tags and Discard)") + print(f"{'='*80}\n") + + # Model + model = ledgers.get_model() + + # Data + dataloader = ledgers.get_dataloader() # Get dataloader + dfm = ledgers.get_dataframe('sample_stats') # Get dataframe manager + + # Optimizer and criterion + optimizer = ledgers.get_optimizer() + criterion = ledgers.get_signal(name="train_mlt_loss/CE") + criterion_bin = ledgers.get_signal(name="train_bin_loss/BCE") + + print("Modifying data...") + + # Add 20 random tags with 'ugly' + tagged_samples = random.sample(range(10), 4) + rows = [] + uids_discarded = [] + for idx in tagged_samples: + uid = dfm._df.index[idx] + uids_discarded.append(uid) + rows.append( + { + "sample_id": uid, + "tags": f"ugly_{random.randint(0, 10)}", # Random tag with 'ugly' + "deny_listed": bool(1 - dfm._df['deny_listed'].iloc[idx]) + } + ) + + # Updates data - Simulate adding tags and discarding samples in dataset + df_update = pd.DataFrame(rows).set_index("sample_id") + # upsert_df updates the ledger's dataframe immediately + dfm.upsert_df(df_update, origin='train_loader', force_flush=True) + + # Changes will be pending + print(f" Added 'ugly' tag to 20 samples") + print(f" Discarded 20 samples") + + # Update hash + exp_hash_d, _, changed = self.chkpt_manager.update_experiment_hash() + + print(f"\n[OK] New experiment hash D: {exp_hash_d}") + print(f"[OK] Changed components: {changed}") + self.assertIn('data', changed, "Data should have changed") + self.assertNotEqual(self.state['exp_hash_c'], exp_hash_d, "Hash should be different") + + print("\nResuming training for 11 epochs...") + pause_controller.resume() # Pending changes to dump: data state + loss_D, uids_D = self.train_epochs( + model, dataloader, optimizer, criterion, + num_epochs=self.config['training']['num_epochs'], + criterion_bin=criterion_bin + ) + pause_controller.pause() + + print("\nTraining completed.") + + # Verify new data directory + data_dir_d = self.chkpt_manager.data_checkpoint_dir / exp_hash_d[-8:] + self.assertTrue(data_dir_d.exists(), "New data checkpoint directory should exist") + + # Verify model weights still being saved + model_dir_d = self.chkpt_manager.models_dir / exp_hash_d[8:-8] + weight_files_d = list(model_dir_d.glob("*_step_*.pt")) + print(f"[OK] Found {len(weight_files_d)} weight checkpoint files") + self.assertGreaterEqual(len(weight_files_d), 2, "Should have at least 2 weight checkpoints") + + # Store state + self.state['exp_hash_d'] = exp_hash_d + self.state['losses_d'] = sum(loss_D) / len(loss_D) + self.state['uids_d'] = uids_D + self.state['uids_discarded_d'] = uids_discarded + self.state['model_c1_neurons'] = model.layers[0].out_neurons + + # Final verbose + print(f"\n[OK] TEST D PASSED - Data state updated") + print(f" Final model_age (i.e., how many epochs lived by the model): {model.current_step}") + + # ======================================================================== + # Test: 05_train_E_reload_and_branch + # ======================================================================== + def test_05_train_E_reload_and_branch(self): + """Reload state B and branch with modified HP and data""" + print(f"\n{'='*80}") + print("TEST E: Reload State B and Branch") + print(f"{'='*80}\n") + + # Get hp from original training + hp_original = self.config + exp_name = hp_original['experiment_name'] + + print("Experiment paused. Analyzing experiment history...") + + # Get all hashes + all_hashes = self.chkpt_manager.get_all_hashes(sort_by='created') + print(f"\n[OK] Found {len(all_hashes)} experiment states:") + for i, entry in enumerate(all_hashes): + print(f" {i+1}. {entry['hash'][:16]}... (created: {entry['created'][:19]})") + + # Reload state B (second state created) + hash_a_from_manifest = self.state['exp_hash_a'] + + print(f"\n[OK] Reloading state B: {hash_a_from_manifest[:16]}...") + + # Use new load_state method to load and apply checkpoint in-place + success = self.chkpt_manager.load_state(exp_hash=hash_a_from_manifest) + self.assertTrue(success, "State should be loaded successfully") + + # Get components from ledger (they were updated in-place) + model_reloaded = ledgers.get_model() + hp_reloaded = ledgers.get_hyperparams(exp_name) + + print(f"[OK] State applied successfully") + print(f"[OK] Loaded HP: {hp_reloaded}") + + # Modify HP and data + print("\nModifying HP and data (not training yet)...") + + # Handle nested config structure + if 'data' in hp_reloaded and 'train_loader' in hp_reloaded['data']: + hp_reloaded['data']['train_loader']['batch_size'] = 1 + old_batch_size = hp_original.get('data', {}).get('train_loader', {}).get('batch_size', 2) + print(f" Batch size: {old_batch_size} -> 1") + + # Discard more data + # Add 20 random tags with 'ugly' + tagged_samples = random.sample(range(10), 1) + rows = [] + dfm = ledgers.get_dataframe() # Get dataframe manager + for idx in tagged_samples: + uid = dfm._df.index[idx] + rows.append( + { + "sample_id": uid, + "tags": f"hugly_{random.randint(0, 10)}", + "deny_listed": bool(1 - dfm._df['deny_listed'].iloc[idx]) + } + ) + # # # Updates data - Simulate adding tags and discarding samples in dataset + # # # upsert_df updates the ledger's dataframe immediately + dfm.upsert_df(pd.DataFrame(rows).set_index("sample_id"), origin='train_loader', force_flush=True) + + # Update hash with all changes + exp_hash_e, _, changed = self.chkpt_manager.update_experiment_hash() + + print(f"\n[OK] New experiment hash E (branch): {exp_hash_e}") + print(f"[OK] Changed components: {changed}") + self.assertIn('hp', changed, "HP should have changed") + self.assertIn('data', changed, "Data should have changed") + + # Update ledger + # Ledger is already registered as proxy are used + pass + + # Setting training environment from loader + dataloader = ledgers.get_dataloader() + model = ledgers.get_model() + optimizer = ledgers.get_optimizer() + criterion = ledgers.get_signal(name="train_mlt_loss/CE") + criterion_bin = ledgers.get_signal(name="train_bin_loss/BCE") + + print("\nResuming training for 21 epochs...") + pause_controller.resume() + loss_E, uids_E = self.train_epochs( + model_reloaded, dataloader, optimizer, criterion, criterion_bin=criterion_bin, + num_epochs=self.config['training']['num_epochs'] * 2, + ) + pause_controller.pause() + + print("\nTraining completed.") + + # Verify checkpoints for E + model_dir_e = self.chkpt_manager.models_dir / exp_hash_e[8:-8] + weight_files_e = list(model_dir_e.glob("*_step_*.pt")) + print(f"[OK] Found {len(weight_files_e)} weight checkpoint files") + self.assertGreaterEqual(len(weight_files_e), 4, "Should have at least 4 weight checkpoints for 21 epochs") + + # Store state + self.state['exp_hash_e'] = exp_hash_e + self.state['losses_e'] = loss_E + self.state['uids_e'] = uids_E + + print(f"\n[OK] TEST E PASSED - Reloaded and generate a new train branch successfully") + print(f" Final model_age: {model.current_step}") + + # ======================================================================== + # Test: 06_reload_before_model_change + # ======================================================================== + def test_06_reload_before_model_change(self): + """Reload before model change (back to A), fix conv size with RNG replay, verify HP+data""" + print(f"\n{'='*80}") + print("TEST 06: Reload Before Model Change - Fix Conv Size with RNG State") + print(f"{'='*80}\n") + + exp_name = self.config['experiment_name'] + hash_A_original = self.state['exp_hash_a'] # Before model change + loss_A_original = self.state['losses_a'] # Before model change + uids_A_original = self.state['uids_a'] # Before model change + + print(f"Reloading state A (before model change) for verification: {hash_A_original[:16]}...") + success = self.chkpt_manager.load_state(exp_hash=hash_A_original) + self.assertTrue(success, "State A should load successfully") + + # Verify HP and data are from checkpoint A + hp_reloaded = ledgers.get_hyperparams(exp_name) + + print(f"[OK] HP batch_size: {hp_reloaded.get('data', {}).get('train_loader', {}).get('batch_size', 'N/A')}") + self.assertEqual(hp_reloaded.get('data', {}).get('train_loader', {}).get('batch_size'), 2, + "Should have batch_size=2 from state A") + print(f"[OK] Data state verified from state A") + print(f"[OK] RNG state restored for reproducible batching") + + # Train with original model to verify batches are the same + print("\nTraining with original model from state A (11 epochs)...") + model_original = ledgers.get_model() + dataloader_original = ledgers.get_dataloader() + optimizer_original = ledgers.get_optimizer() + criterion = ledgers.get_signal(name="train_mlt_loss/CE") + criterion_bin = ledgers.get_signal(name="train_bin_loss/BCE") + + pause_controller.resume() + loss_A_reloaded, uids_A_reloaded = self.train_epochs( + model_original, dataloader_original, optimizer_original, criterion, + num_epochs=self.config['training']['num_epochs'], + criterion_bin=criterion_bin + ) + pause_controller.pause() + + # Check reproducibility with original loss and UIDs + self.check_reproducibility(loss_A_original, loss_A_reloaded, uids_A_original, uids_A_reloaded) + + # Reload again and fix model, should get same batches due to restored RNG + print(f"\nReloading state A again (to reset RNG for fair comparison) and modifying model architecture...") + success = self.chkpt_manager.load_state(exp_hash=hash_A_original) + self.assertTrue(success, "State A should load successfully second time") + + # Fix model conv size - create new model with different architecture + print("\nFixing model architecture...") + model = ledgers.get_model() + model.operate(0, {-1}, 1) + model.operate(2, {-1}, 2) + model.operate(-2, {}, 3) + model.operate(-1, {-1 }, 4) + + exp_hash_h, _, changed = self.chkpt_manager.update_experiment_hash() + print(f"\n[OK] New experiment hash H: {exp_hash_h[:16]}") + print(f"[OK] Changed components: {changed}") + self.assertIn('model', changed, "Only model should have changed") + self.assertNotIn('hp', changed, "HP should not have changed") + self.assertNotIn('data', changed, "Data should not have changed") + + # Train with new model - should get same batches due to restored RNG + dataloader = ledgers.get_dataloader() + optimizer = ledgers.get_optimizer() + criterion = ledgers.get_signal(name="train_mlt_loss/CE") + criterion_bin = ledgers.get_signal(name="train_bin_loss/BCE") + + print("\nTraining for 11 epochs with new model (same RNG state = same batches)...") + pause_controller.resume() + loss_H, uids_H = self.train_epochs(model, dataloader, optimizer, criterion, num_epochs=self.config['training']['num_epochs'], + criterion_bin=criterion_bin) + pause_controller.pause() + + print(f"[OK] Fixed model training loss (first/last): {loss_H} / {loss_H}") + + # Compare: First batch should be same, but losses differ due to different model + print(f"\n[OK] Reproducibility verified:") + print(f" Original model first batch loss: {loss_A_reloaded}") + print(f" Fixed model first batch loss: {loss_H}") + print(f" (Same RNG = same batches, different losses due to model change)") + + # Store state + self.state['losses_h'] = loss_H + self.state['exp_hash_h'] = exp_hash_h + self.state['uids_h'] = uids_H + + print(f"\n[OK] TEST 06 PASSED - Reloaded with RNG state, trained with fixed architecture") + + # ======================================================================== + # Test: 07_change_data_from_test06 + # ======================================================================== + def test_07_change_data_from_test06(self): + """Change data from test 06 - discard more data and train again""" + print(f"\n{'='*80}") + print("TEST 07: Change Data from Test 06 - Discard More Data") + print(f"{'='*80}\n") + + exp_name = self.config['experiment_name'] + hash_H = self.state['exp_hash_h'] # From test 06 + + print(f"Starting from state H: {hash_H[:16]}...") + + # Discard additional 15 samples (total 25% discarded) + print("\nDiscarding additional 15 samples (25% total)...") + dfm = ledgers.get_dataframe('sample_stats') + tagged_samples = random.sample(range(10), 2) + rows = [] + for idx in tagged_samples: + uid = dfm._df.index[idx] + rows.append({ + "sample_id": uid, + "tags": f"discard_25pct_{random.randint(0, 10)}", + "deny_listed": True + }) + + df_update = pd.DataFrame(rows).set_index("sample_id") + dfm.upsert_df(df_update, origin='train_loader', force_flush=True) + + exp_hash_i, _, changed = self.chkpt_manager.update_experiment_hash() + print(f"\n[OK] New experiment hash I: {exp_hash_i[:16]}") + print(f"[OK] Changed components: {changed}") + self.assertIn('data', changed, "Only data should have changed") + + # Train for 11 epochs + model = ledgers.get_model() + dataloader = ledgers.get_dataloader() + optimizer = ledgers.get_optimizer() + criterion = ledgers.get_signal(name="train_mlt_loss/CE") + criterion_bin = ledgers.get_signal(name="train_bin_loss/BCE") + + print("\nTraining for 11 epochs with 25% discarded...") + pause_controller.resume() + loss_I, uids_I = self.train_epochs(model, dataloader, optimizer, criterion, num_epochs=self.config['training']['num_epochs'], + criterion_bin=criterion_bin) + pause_controller.pause() + + # Store state + self.state['losses_i'] = loss_I + self.state['uids_i'] = uids_I + self.state['exp_hash_i'] = exp_hash_i + + print(f"\n[OK] TEST 07 PASSED - Changed data and trained successfully") + + # ======================================================================== + # Test: 08_reload_before_data_change_verify_and_modify + # ======================================================================== + def test_08_reload_before_data_change_verify_and_modify(self): + """Reload before data change (state C), verify training reproducibility, then modify model""" + print(f"\n{'='*80}") + print("TEST 08: Reload Before Data Change - Verify and Modify Model") + print(f"{'='*80}\n") + + exp_name = self.config['experiment_name'] + hash_c = self.state['exp_hash_c'] # Before data change (after HP change) + loss_c = self.state['losses_c'] + uids_c = self.state.get('uids_c') + + print(f"Part A: Reloading state C and verifying training reproducibility...") + print(f"Reloading state C: {hash_c[:16]}...") + + success = self.chkpt_manager.load_state(exp_hash=hash_c) + self.assertTrue(success, "State C should load successfully") + + # Verify training produces same results + model = ledgers.get_model() + dataloader = ledgers.get_dataloader() + optimizer = ledgers.get_optimizer() + criterion = ledgers.get_signal(name="train_mlt_loss/CE") + criterion_bin = ledgers.get_signal(name="train_bin_loss/BCE") + + print("\nTraining for 11 epochs to verify reproducibility...") + pause_controller.resume() + loss_C_verify, uids_C_verify = self.train_epochs(model, dataloader, optimizer, criterion, num_epochs=self.config['training']['num_epochs'], + criterion_bin=criterion_bin) + pause_controller.pause() + + # Check reproducibility with original loss and UIDs + # self.check_reproducibility(loss_c, loss_C_verify, uids_c, None, loss_tol=1e-1) + + print(f"\nPart B: Modifying model from reloaded state C...") + + # Reload again to reset state + success = self.chkpt_manager.load_state(exp_hash=hash_c) + + # Modify model + model = ledgers.get_model() + print("\nModifying model architecture...") + model.operate(0, {-2}, 1) # Change conv1 + model.operate(2, {-2}, 2) # Change conv2 + + exp_hash_j, _, changed = self.chkpt_manager.update_experiment_hash() + print(f"\n[OK] New experiment hash J: {exp_hash_j[:16]}") + self.assertIn('model', changed, "Model should have changed") + + # Train with modified model + dataloader = ledgers.get_dataloader() + optimizer = ledgers.get_optimizer() + + print("\nTraining for 11 epochs with modified model...") + pause_controller.resume() + loss_J, _ = self.train_epochs(model, dataloader, optimizer, criterion, num_epochs=self.config['training']['num_epochs'], + criterion_bin=criterion_bin) + pause_controller.pause() + + # Store state + self.state['losses_j'] = sum(loss_J)/len(loss_J) + self.state['exp_hash_j'] = exp_hash_j + + print(f"\n[OK] TEST 08 PASSED - Verified reproducibility and modified model") + + # ======================================================================== + # Test: 09_reload_before_hp_change_verify_and_fix + # ======================================================================== + def test_09_reload_before_hp_change_verify_and_modify(self): + """Reload before HP change (state B), verify training, then fix HP, model, and data""" + print(f"\n{'='*80}") + print("TEST 09: Reload Before HP Change - Verify and Fix Everything") + print(f"{'='*80}\n") + + exp_name = self.config['experiment_name'] + hash_b = self.state['exp_hash_b'] # Before HP change (after model change) + loss_b = self.state['losses_b'] + + print(f"Part A: Reloading state B and verifying training reproducibility...") + print(f"Reloading state B: {hash_b[:16]}...") + + success = self.chkpt_manager.load_state(exp_hash=hash_b) + self.assertTrue(success, "State B should load successfully") + + # Verify training produces same results + model = ledgers.get_model() + dataloader = ledgers.get_dataloader() + optimizer = ledgers.get_optimizer() + criterion = ledgers.get_signal(name="train_mlt_loss/CE") + criterion_bin = ledgers.get_signal(name="train_bin_loss/BCE") + + print("\nTraining for 11 epochs to verify reproducibility...") + pause_controller.resume() + loss_B_verify, uids_B_verify = self.train_epochs(model, dataloader, optimizer, criterion, num_epochs=self.config['training']['num_epochs'], + criterion_bin=criterion_bin) + pause_controller.pause() + + # Check reproducibility with original loss and UIDs + self.check_reproducibility(loss_b, loss_B_verify, self.state.get('uids_b'), None, loss_tol=1e-1) + + print(f"\nPart B: Fixing HP, model, and data from reloaded state B...") + + # Reload again to reset state + success = self.chkpt_manager.load_state(exp_hash=hash_b) + + # Fix HP + hp = ledgers.get_hyperparams() + hp['data']['train_loader']['batch_size'] = 7 # Change batch size + + # Fix model + model = ledgers.get_model() + model.operate(0, {-3}, 1) # Further modify conv1 + + # Fix data - discard 5 samples + dfm = ledgers.get_dataframe('sample_stats') + tagged_samples = random.sample(range(10), 2) + rows = [] + for idx in tagged_samples: + uid = dfm._df.index[idx] + rows.append({ + "sample_id": uid, + "tags": f"discard_fix_{random.randint(0, 10)}", + "deny_listed": True + }) + df_update = pd.DataFrame(rows).set_index("sample_id") + dfm.upsert_df(df_update, origin='train_loader', force_flush=True) + + exp_hash_k, _, changed = self.chkpt_manager.update_experiment_hash() + print(f"\n[OK] New experiment hash K: {exp_hash_k[:16]}") + print(f"[OK] Changed components: {changed}") + self.assertIn('hp', changed, "HP should have changed") + self.assertIn('model', changed, "Model should have changed") + self.assertIn('data', changed, "Data should have changed") + + # Train with all fixes + dataloader = ledgers.get_dataloader() + optimizer = ledgers.get_optimizer() + criterion = ledgers.get_signal(name="train_mlt_loss/CE") + criterion_bin = ledgers.get_signal(name="train_bin_loss/BCE") + + print("\nTraining for 11 epochs with all fixes...") + pause_controller.resume() + loss_K, _ = self.train_epochs(model, dataloader, optimizer, criterion, num_epochs=self.config['training']['num_epochs'], + criterion_bin=criterion_bin) + pause_controller.pause() + + # Store state + self.state['losses_k'] = sum(loss_K)/len(loss_K) + self.state['exp_hash_k'] = exp_hash_k + + print(f"\n[OK] TEST 09 PASSED - Verified reproducibility and fixed everything") + + # ======================================================================== + # Test: 10_reload_branch_j_verify_reproducibility + # ======================================================================== + def test_10_reload_branch_j_verify_reproducibility(self): + """Reload branch J (from test 08.b) and verify training reproducibility""" + print(f"\n{'='*80}") + print("TEST 10: Reload Branch J - Verify Training Reproducibility") + print(f"{'='*80}\n") + + exp_name = self.config['experiment_name'] + hash_j = self.state['exp_hash_j'] # From test 08.b + loss_j = self.state['losses_j'] + + print(f"Reloading branch J: {hash_j[:16]}...") + + success = self.chkpt_manager.load_state(exp_hash=hash_j) + self.assertTrue(success, "State J should load successfully") + + # Train again to verify reproducibility + model = ledgers.get_model() + dataloader = ledgers.get_dataloader() + optimizer = ledgers.get_optimizer() + criterion = ledgers.get_signal(name="train_mlt_loss/CE") + criterion_bin = ledgers.get_signal(name="train_bin_loss/BCE") + + print("\nTraining for 11 epochs to verify reproducibility...") + pause_controller.resume() + loss_j_verify, _ = self.train_epochs(model, dataloader, optimizer, criterion, num_epochs=self.config['training']['num_epochs'], + criterion_bin=criterion_bin) + pause_controller.pause() + + # Check reproducibility with original loss and UIDs + self.check_reproducibility(self.state['losses_j'], loss_j_verify, self.state.get('uids_b'), None, loss_tol=1e-1) + + print(f"\n[OK] TEST 10 PASSED - Branch J training is reproducible") + + # ======================================================================== + # Test: 11_restart_from_config_verify_reproducibility + # ======================================================================== + def test_11_restart_from_scratch_to_hash_d_and_verify_reproducibility(self): + """Test 11: Restart experiment from config - verify all components load to branch_j state""" + print(f"\n{'='*80}") + print("TEST 11: Restart Experiment from Config - Verify Full Reproducibility") + print(f"{'='*80}\n") + + # Reference variables + target_hash = self.state['exp_hash_d'] # Target is branch_d + loss_d_original = self.state['losses_d'] + originals_uids = self.state.get('uids_d', None) + + print(f"Simulating fresh restart: loading everything from config...") + print(f"Target state: {target_hash[:16]} (branch_d)") + + # Simulate fresh Python process: re-register everything from config + config_reloaded = self.config_cp + exp_name = EXP_NAME + + # Clear existing ledger entries + ledgers.clear_all() + print("[OK] Cleared existing ledger entries") + + # ================================================= + # Automotically load components from existing chkpt + # ================================================= + # First init a checkpoint manager with reloaded config + self.chkpt_manager = CheckpointManagerV2(root_log_dir=self.config.get('root_log_dir')) + ledgers.register_checkpoint_manager(exp_name, self.chkpt_manager) + + # Re-register HP + register_in_ledger(config_reloaded, flag="hyperparameters", name=exp_name) + print("[OK] Hyperparameters re-registered") + + # Create fresh model + model_restarted = SimpleCNN(conv1_out=8, conv2_out=16) # Match branch_d architecture + # # Model arch. and weights are updated at the init of model interface + model_restarted = register_in_ledger(model_restarted, flag="model", name=exp_name, device=DEVICE) + + # Re-register dataloader + # # Same here, dataloader is created from HP at init of dataloader interface, and data are loaded from chkpt + register_in_ledger( + self.dataset, + flag="dataloader", + name=exp_name, + compute_hash=False, + is_training=True, + batch_size=config_reloaded.get('data', {}).get('train_loader', {}).get('batch_size', 32), + shuffle=config_reloaded.get('data', {}).get('train_loader', {}).get('shuffle', False) + ) + + # Create and register dataloader + dataloader = register_in_ledger( + self.dataset, + flag="dataloader", + name=exp_name, + compute_hash=False, + is_training=True, + batch_size=self.config.get('data', {}).get('train_loader', {}).get('batch_size', 32), + shuffle=self.config.get('data', {}).get('train_loader', {}).get('shuffle', False) + ) + + # Optimizer and criterion + optimizer_restarted = th.optim.Adam( + model_restarted.parameters(), + lr=config_reloaded.get('optimizer', {}).get('lr', 0.001) + ) + optimizer_restarted = register_in_ledger(optimizer_restarted, flag="optimizer", name=exp_name) + # # Create and register signal (criterion) + criterion = nn.CrossEntropyLoss(reduction='none') + criterion = register_in_ledger( + criterion, + flag="signal", + name="train_mlt_loss/CE", + log=True + ) + criterion_bin = nn.BCEWithLogitsLoss(reduction='none') + criterion_bin = register_in_ledger( + criterion_bin, + flag="signal", + name="train_bin_loss/BCE", + log=True + ) + print("[OK] Fresh registrations complete") + + # Get all hashes + all_hashes = self.chkpt_manager.get_all_hashes(sort_by='created') + print(f"\n[OK] Found {len(all_hashes)} experiment states:") + for i, entry in enumerate(all_hashes): + print(f" {i+1}. {entry['hash'][:16]}... (created: {entry['created'][:19]})") + + # Reload state B (second state created) + hash_a_from_manifest = self.state['exp_hash_a'] + + print(f"\n[OK] Reloading state B: {hash_a_from_manifest[:16]}...") + + # Use new load_state method to load and apply checkpoint in-place + success = self.chkpt_manager.load_state(exp_hash=target_hash) + self.assertTrue(success, "State should be loaded successfully") + + print(f"[OK] Checkpoint loaded to reach target state {target_hash[:16]}") + print("\nTraining for 11 epochs to verify reproducibility...") + pause_controller.resume() + _, _ = self.train_epochs(model_restarted, dataloader, optimizer_restarted, criterion, num_epochs=self.config['training']['num_epochs'], + criterion_bin=criterion_bin) + pause_controller.pause() + + # Check reproducibility with original loss and UIDs + self.assertEqual(model_restarted.layers[-1].operation_age['FREEZE'], 1, + "Model architecture should match state in D") + self.assertEqual(model_restarted.layers[-1].operation_age['RESET'], 1, + "Model architecture should match state in D") + self.assertEqual(model_restarted.layers[0].out_neurons, 8, + "Model architecture should match state in D") + + # Not possible as data are generated randomly without reproducibility now + # self.check_reproducibility(loss_d_original, loss_d_verify, originals_uids, None, loss_tol=1e-1) + + # ======================================================================== + # Test: logger queue saved with weights + # ======================================================================== + def test_logger_queue_saved_with_weights(self): + self.chkpt_manager.update_experiment_hash(force=False, dump_immediately=False) + + snapshot_path = Path(self.chkpt_manager.loggers_dir) / self.chkpt_manager.current_exp_hash / "loggers.json" + self.assertTrue(snapshot_path.exists(), "Logger snapshot should be saved with checkpoint") + + with open(snapshot_path, "r") as f: + snapshot = json.load(f) + + loggers = snapshot.get("loggers", {}) + self.assertIn(self.config.get("experiment_name"), loggers, "Logger entry should be present") + signals = loggers[self.config.get("experiment_name")].get("signal_history", []) + self.assertGreaterEqual(len(signals), 1, "Signal history should contain logged signals") + + +if __name__ == '__main__': + # Create test suite with explicit ordering + suite = unittest.TestSuite() + + # Add tests in specific order + # # Initialize experiment + suite.addTest(CheckpointSystemV3Tests('test_00_initialize_experiment')) + # # User Adventures training workflow + suite.addTest(CheckpointSystemV3Tests('test_01_train_A')) + suite.addTest(CheckpointSystemV3Tests('test_02_train_B_model_change')) + suite.addTest(CheckpointSystemV3Tests('test_03_train_C_hyperparams_change')) + suite.addTest(CheckpointSystemV3Tests('test_04_train_D_data_change')) + # # Reload and branching tests + suite.addTest(CheckpointSystemV3Tests('test_05_train_E_reload_and_branch')) + suite.addTest(CheckpointSystemV3Tests('test_06_reload_before_model_change')) + suite.addTest(CheckpointSystemV3Tests('test_07_change_data_from_test06')) + # # Reload and check full reproducibility - Loss and UIDs + suite.addTest(CheckpointSystemV3Tests('test_08_reload_before_data_change_verify_and_modify')) + suite.addTest(CheckpointSystemV3Tests('test_09_reload_before_hp_change_verify_and_modify')) + suite.addTest(CheckpointSystemV3Tests('test_10_reload_branch_j_verify_reproducibility')) + suite.addTest(CheckpointSystemV3Tests('test_11_restart_from_scratch_to_hash_d_and_verify_reproducibility')) + # # Check that logger queue is saved and loaded + suite.addTest(CheckpointSystemV3Tests('test_logger_queue_saved_with_weights')) + + # Run the suite + runner = unittest.TextTestRunner(verbosity=2) + runner.run(suite) diff --git a/weightslab/tests/test_data_loader_interface.py b/weightslab/tests/test_data_loader_interface.py index 7fd536e4..03a9820a 100644 --- a/weightslab/tests/test_data_loader_interface.py +++ b/weightslab/tests/test_data_loader_interface.py @@ -1,7 +1,11 @@ import math import unittest import torch -from torch.utils.data import TensorDataset, DataLoader +from torch.utils.data import TensorDataset, DataLoader, Subset +from torchvision import datasets, transforms + +from weightslab.backend.dataloader_interface import DataLoaderInterface +from weightslab.utils.tools import capture_rng_state, restore_rng_state, seed_everything def infinite_loader(loader): @@ -76,6 +80,154 @@ def test_infinite_loader_restarts_epochs_and_collects_all_labels(self): self.assertEqual(len(set(labels)), len(self.train_loader.dataset)) -if __name__ == "__main__": - unittest.main() +class TestDataLoaderReproducibility(unittest.TestCase): + """Test RNG and iteration state reproducibility for dataloaders.""" + + @classmethod + def setUpClass(cls): + """Set up test dataset once for all reproducibility tests.""" + # Use a small MNIST subset for testing + transform = transforms.Compose([ + transforms.ToTensor(), + transforms.Normalize((0.1307,), (0.3081,)) + ]) + + try: + # Try to load from common location + full_dataset = datasets.MNIST( + root='C:/Users/GuillaumePelluet/Desktop/mnist_data/', + train=False, + download=False, + transform=transform + ) + except: + # Fallback to temp directory + import tempfile + temp_dir = tempfile.mkdtemp() + full_dataset = datasets.MNIST( + root=temp_dir, + train=False, + download=True, + transform=transform + ) + + # Create subset with 100 samples + subset_indices = list(range(100)) + cls.dataset = Subset(full_dataset, subset_indices) + + def test_rng_reproducibility_with_shuffle(self): + """Test dataloader reproducibility with shuffle: save RNG → generate batches → reload RNG → verify same batches. + + Key insight: Shuffle happens when iter() is called. Restoring RNG before + reset_iterator() ensures identical shuffle ordering. + """ + print(f"\n{'='*60}") + print("RNG State Reproducibility - Shuffle Enabled") + print(f"{'='*60}\n") + + # 1. Initialize with seed and create dataloader + print("1. Initializing with seed=42...") + seed_everything(42) + + dataloader = DataLoaderInterface( + self.dataset, + batch_size=2, + shuffle=True, + num_workers=0 + ) + print(f"[OK] DataLoader created (batch_size=2, shuffle=True)") + + # Consume initial batches + _, bids_1_init = next(dataloader) + _, bids_2_init = next(dataloader) + print(f"Initial warmup batches: {bids_1_init.tolist()}, {bids_2_init.tolist()}") + + # 2. Capture RNG state + print("\n2. Capturing RNG state...") + rng_state = capture_rng_state() + dataloader.reset_iterator() # Reset to use captured RNG + print(f"[OK] RNG state captured and iterator reset") + + # 3. Generate batches with current RNG + print("\n3. Generating batches...") + _, bids_1 = next(dataloader) + _, bids_2 = next(dataloader) + print(f"Batches: {bids_1.tolist()}, {bids_2.tolist()}") + + # 4. Restore RNG and reset iterator + print("\n4. Restoring RNG state and resetting iterator...") + restore_rng_state(rng_state) + dataloader.reset_iterator() + print(f"[OK] RNG restored, iterator reset") + + # 5. Generate batches again - should be identical + print("\n5. Generating batches with restored RNG...") + _, bids_1_repeat = next(dataloader) + _, bids_2_repeat = next(dataloader) + print(f"Repeated batches: {bids_1_repeat.tolist()}, {bids_2_repeat.tolist()}") + + # Verify + print(f"\n{'='*60}") + print("Verification:") + print(f" Batch 1 match: {torch.equal(bids_1, bids_1_repeat)}") + print(f" Batch 2 match: {torch.equal(bids_2, bids_2_repeat)}") + self.assertTrue(torch.equal(bids_1, bids_1_repeat), "First batches should be identical") + self.assertTrue(torch.equal(bids_2, bids_2_repeat), "Second batches should be identical") + print(f"[OK] RNG reproducibility verified!\n") + + def test_iteration_state_reproducibility_without_shuffle(self): + """Test dataloader reproducibility without shuffle: capture iteration state → resume identically. + + With shuffle disabled, RNG is irrelevant. We capture the iteration position + (number of batches yielded) and restore that position efficiently using + OffsetSampler to skip samples at the index level without data reprocessing. + """ + print(f"\n{'='*60}") + print("Iteration State Reproducibility - No Shuffle") + print(f"{'='*60}\n") + + print("1. Creating dataloader (shuffle=False)...") + dataloader = DataLoaderInterface( + self.dataset, + batch_size=2, + shuffle=False, + num_workers=0 + ) + print(f"[OK] DataLoader created (batch_size=2, shuffle=False)") + + # 2. Consume two batches, then capture state + print("\n2. Consuming first 2 batches...") + _, bids_1 = next(dataloader) + _, bids_2 = next(dataloader) + print(f"Batches 1-2: {bids_1.tolist()}, {bids_2.tolist()}") + + iter_state = dataloader.capture_iteration_state() + print(f"[OK] Iteration state captured: {iter_state}") + + # 3. Consume next two batches + print("\n3. Consuming batches 3-4...") + _, bids_3 = next(dataloader) + _, bids_4 = next(dataloader) + print(f"Batches 3-4: {bids_3.tolist()}, {bids_4.tolist()}") + + # 4. Restore iteration state + print(f"\n4. Restoring to position after batch 2...") + dataloader.restore_iteration_state(iter_state) + print(f"[OK] Iteration state restored (skipped first 2 batches efficiently)") + + # 5. Generate batches again - should match 3 and 4 + print("\n5. Generating next batches (should match 3-4)...") + _, bids_3_repeat = next(dataloader) + _, bids_4_repeat = next(dataloader) + print(f"Repeated batches: {bids_3_repeat.tolist()}, {bids_4_repeat.tolist()}") + + # Verify + print(f"\n{'='*60}") + print("Verification:") + print(f" Batch 3 match: {torch.equal(bids_3, bids_3_repeat)}") + print(f" Batch 4 match: {torch.equal(bids_4, bids_4_repeat)}") + self.assertTrue(torch.equal(bids_3, bids_3_repeat), "Batch 3 should be identical") + self.assertTrue(torch.equal(bids_4, bids_4_repeat), "Batch 4 should be identical") + print(f"[OK] Iteration state reproducibility verified!\n") + diff --git a/weightslab/trainer/experiment_context.py b/weightslab/trainer/experiment_context.py index b978b371..31d8408f 100644 --- a/weightslab/trainer/experiment_context.py +++ b/weightslab/trainer/experiment_context.py @@ -1,5 +1,9 @@ import logging +from weightslab.components.global_monitoring import pause_controller + + +# Init global logger logger = logging.getLogger(__name__) @@ -36,6 +40,8 @@ def ensure_components(self): logger). Raises RuntimeError when mandatory components are missing. """ from weightslab.backend.ledgers import ( + get_checkpoint_manager, + list_checkpoint_managers, get_hyperparams, list_hyperparams, get_model, @@ -107,11 +113,24 @@ def ensure_components(self): except Exception: signal_logger = None + # resolve checkpoint manager + checkpoint_manager = None + try: + lnames = list_checkpoint_managers() + if len(lnames) == 1: + checkpoint_manager = get_checkpoint_manager() + elif "main" in lnames: + checkpoint_manager = get_checkpoint_manager("main") + except Exception: + checkpoint_manager = None + self._components = { "model": model, "optimizer": optimizer, "hyperparams": hyperparams, "signal_logger": signal_logger, + "trainer": pause_controller, + "checkpoint_manager": checkpoint_manager } self._components.update(data_loaders) # add all dataloaders found @@ -148,15 +167,15 @@ def _get_total_steps(): current = int(model.current_step) elif hasattr(model, 'get_age'): current = int(model.get_age()) - + # Get remaining from hyperparams remaining = _hp_getter("training_steps_to_do", 999)() - + # If explicit total is set, use it. Otherwise calculate. explicit_total = _hp_getter("total_training_steps", None)() if explicit_total is not None: return explicit_total - + return current + int(remaining) except Exception: return 1000 diff --git a/weightslab/trainer/services/data_service.py b/weightslab/trainer/services/data_service.py index fe67bf84..34329a04 100644 --- a/weightslab/trainer/services/data_service.py +++ b/weightslab/trainer/services/data_service.py @@ -1274,6 +1274,8 @@ def ApplyDataQuery(self, request, context): - number_of_samples_in_the_loop: rows not deny_listed - number_of_discarded_samples: rows with deny_listed == True """ + self._ctx.ensure_components() + components = self._ctx.components # 1) No query: just report counts (Needs lock for consistency) if request.query == "": @@ -1285,6 +1287,17 @@ def ApplyDataQuery(self, request, context): ) try: + # Pause training if it's currently running + trainer = components.get("trainer") + hp = components.get("hyperparams") + if trainer: + logger.info("Pausing training before restore...") + trainer.pause() + if "is_training" in hp: + hp['is_training'] = False + else: + hp["is_training"] = False + # 2) Check if we should bypass the agent (Quick Filters path) if not request.is_natural_language: logger.info( @@ -1492,6 +1505,18 @@ def EditDataSample(self, request, context): self._initialize_data_service() self._ctx.ensure_components() + components = self._ctx.components + + # Pause training if it's currently running + trainer = components.get("trainer") + hp = components.get("hyperparams") + if trainer: + logger.info("Pausing training before restore...") + trainer.pause() + if "is_training" in hp: + hp['is_training'] = False + else: + hp["is_training"] = False if request.stat_name not in [SampleStatsEx.TAGS.value, SampleStatsEx.DENY_LISTED.value]: return pb2.DataEditsResponse( diff --git a/weightslab/trainer/services/experiment_service.py b/weightslab/trainer/services/experiment_service.py index 6b589488..ae81e8a9 100644 --- a/weightslab/trainer/services/experiment_service.py +++ b/weightslab/trainer/services/experiment_service.py @@ -1,6 +1,7 @@ import time -import logging import types +import queue +import logging import weightslab.proto.experiment_service_pb2 as pb2 from weightslab.components.global_monitoring import weightslab_rlock @@ -8,6 +9,8 @@ from weightslab.trainer.services.model_service import ModelService from weightslab.trainer.services.data_service import DataService + +# Logger logger = logging.getLogger(__name__) @@ -31,68 +34,66 @@ def get_root_log_dir(self) -> str: return self.data_service.get_root_log_dir() # ------------------------------------------------------------------------- - # Training status stream + # Logger queue sync for WeightsStudio # ------------------------------------------------------------------------- - def stream_status(self, request_iterator): - import queue - + def GetLatestLoggerData(self, request, context): + """ + Returns logger data for WeightsStudio polling. + - If request_full_history is True: returns full history (limited by max_points per signal) + - If request_full_history is False: returns only new data from the queue since last request + """ self._ctx.ensure_components() components = self._ctx.components - signal_logger = components.get("signal_logger") if getattr(self._ctx, "_components", None) else None - - while True: - try: - if signal_logger == None: - # No signal logger available, wait briefly and continue - time.sleep(0.01) - continue - - # Use timeout to avoid blocking indefinitely - try: - signal_log = signal_logger.queue.get(timeout=0.5) - except queue.Empty: - # No signals available, continue waiting - continue - - if "metric_name" in signal_log and "acc" in signal_log["metric_name"]: - logger.debug(f"[signal_log] {signal_log['metric_name']} = {signal_log['metric_value']:.2f}") - - metrics_status, annotat_status = None, None - if "metric_name" in signal_log: - metrics_status = pb2.MetricsStatus( - name=signal_log["metric_name"], - value=signal_log["metric_value"], - ) - elif "annotation" in signal_log: - annotat_status = pb2.AnnotatStatus(name=signal_log["annotation"]) - for key, value in signal_log["metadata"].items(): - annotat_status.metadata[key] = value - - training_status = pb2.TrainingStatusEx( - timestamp=time.strftime("%Y-%m-%d %H:%M:%S"), - experiment_name=signal_log["experiment_name"], - model_age=signal_log["model_age"], - ) - - if metrics_status: - training_status.metrics_status.CopyFrom(metrics_status) - if annotat_status: - training_status.annotat_status.CopyFrom(annotat_status) - - # mark task done on ledger logger queue - try: - signal_logger.queue.task_done() - except Exception: - pass - - yield training_status - - except GeneratorExit: - # Client disconnected, exit gracefully - logger.debug("Stream status client disconnected") - break + signal_logger = components.get("signal_logger") + if signal_logger is None: + return pb2.GetLatestLoggerDataResponse(points=[]) + + points = [] + + if request.request_full_history: + # Return full history + max_points = request.max_points or 10000 + history = signal_logger.get_signal_history() + + # Group by metric_name and limit each + signal_groups = {} + for s in history: + metric_name = s.get("metric_name", "") + if metric_name not in signal_groups: + signal_groups[metric_name] = [] + signal_groups[metric_name].append(s) + + # Take last max_points_per_signal for each signal and downsample if needed + for metric_name, signal_history in signal_groups.items(): + + # Downsample if we have more than 1000 points + if len(signal_history) > max_points: + # Calculate step to downsample (e.g., if 5000 points, step=5 to get ~1000) + step = max(1, len(signal_history) // max_points) + signal_history = signal_history[::step] + + for s in signal_history: + points.append(pb2.LoggerDataPoint( + metric_name=metric_name, + model_age=s.get("model_age", 0), + metric_value=s.get("metric_value", 0.0), + experiment_hash=s.get("experiment_hash", ""), + timestamp=int(s.get("timestamp", time.time())), + )) + else: + # Return only queue (new data since last poll) + queue_data = signal_logger.get_and_clear_queue() + for s in queue_data: + points.append(pb2.LoggerDataPoint( + metric_name=s.get("metric_name", ""), + model_age=s.get("model_age", 0), + metric_value=s.get("metric_value", 0.0), + experiment_hash=s.get("experiment_hash", ""), + timestamp=int(s.get("timestamp", time.time())), + )) + + return pb2.GetLatestLoggerDataResponse(points=points) - # ------------------------------------------------------------------------- # Training & hyperparameter commands # ------------------------------------------------------------------------- def ExperimentCommand(self, request, context): @@ -102,6 +103,17 @@ def ExperimentCommand(self, request, context): # Write requests if request.HasField("hyper_parameter_change"): with weightslab_rlock: + # Pause training if it's currently running + trainer = components.get("trainer") + hp = components.get("hyperparams") + if trainer: + logger.info("Pausing training before restore...") + trainer.pause() + if "is_training" in hp: + hp['is_training'] = False + else: + hp["is_training"] = False + hyper_parameters = request.hyper_parameter_change.hyper_parameters from weightslab.backend.ledgers import set_hyperparam, list_hyperparams, resolve_hp_name hp_name = None @@ -163,6 +175,17 @@ def ExperimentCommand(self, request, context): with weightslab_rlock: from weightslab.backend.ledgers import get_dataloaders + # Pause training if it's currently running + trainer = components.get("trainer") + hp = components.get("hyperparams") + if trainer: + logger.info("Pausing training before restore...") + trainer.pause() + if "is_training" in hp: + hp['is_training'] = False + else: + hp["is_training"] = False + denied_cnt = len(request.deny_samples_operation.sample_ids) origin = request.deny_samples_operation.origin if hasattr(request.deny_samples_operation, 'origin') else 'train' @@ -198,6 +221,17 @@ def ExperimentCommand(self, request, context): with weightslab_rlock: from weightslab.backend.ledgers import get_dataloaders + # Pause training if it's currently running + trainer = components.get("trainer") + hp = components.get("hyperparams") + if trainer: + logger.info("Pausing training before restore...") + trainer.pause() + if "is_training" in hp: + hp['is_training'] = False + else: + hp["is_training"] = False + denied_cnt = len(request.deny_eval_samples_operation.sample_ids) origin = request.deny_eval_samples_operation.origin if hasattr(request.deny_eval_samples_operation, 'origin') else 'eval' @@ -233,6 +267,17 @@ def ExperimentCommand(self, request, context): with weightslab_rlock: from weightslab.backend.ledgers import get_dataloaders + # Pause training if it's currently running + trainer = components.get("trainer") + hp = components.get("hyperparams") + if trainer: + logger.info("Pausing training before restore...") + trainer.pause() + if "is_training" in hp: + hp['is_training'] = False + else: + hp["is_training"] = False + allowed = set(request.remove_from_denylist_operation.sample_ids) origin = request.remove_from_denylist_operation.origin if hasattr(request.remove_from_denylist_operation, 'origin') else 'train' @@ -265,6 +310,17 @@ def ExperimentCommand(self, request, context): with weightslab_rlock: from weightslab.backend.ledgers import get_dataloaders + # Pause training if it's currently running + trainer = components.get("trainer") + hp = components.get("hyperparams") + if trainer: + logger.info("Pausing training before restore...") + trainer.pause() + if "is_training" in hp: + hp['is_training'] = False + else: + hp["is_training"] = False + allowed = set(request.remove_eval_from_denylist_operation.sample_ids) origin = request.remove_eval_from_denylist_operation.origin if hasattr(request.remove_eval_from_denylist_operation, 'origin') else 'eval' @@ -295,6 +351,18 @@ def ExperimentCommand(self, request, context): if request.HasField("load_checkpoint_operation"): with weightslab_rlock: + + # Pause training if it's currently running + trainer = components.get("trainer") + hp = components.get("hyperparams") + if trainer: + logger.info("Pausing training before restore...") + trainer.pause() + if "is_training" in hp: + hp['is_training'] = False + else: + hp["is_training"] = False + checkpoint_id = request.load_checkpoint_operation.checkpoint_id model = components.get("model") if model is None: @@ -373,3 +441,58 @@ def ExperimentCommand(self, request, context): response.sample_statistics.origin = request.get_data_records return response + + def RestoreCheckpoint(self, request, context): + """ + Restore a checkpoint from a given experiment hash. + - Pauses training if not already paused + - Calls checkpoint manager to load the state + - Returns success flag and message + """ + try: + experiment_hash = request.experiment_hash + logger.info(f"Restoring checkpoint from hash: {experiment_hash}") + + self._ctx.ensure_components() + components = self._ctx.components + + # Pause training if it's currently running + trainer = components.get("trainer") + hp = components.get("hyperparams") + if trainer: + logger.info("Pausing training before restore...") + trainer.pause() + if "is_training" in hp: + hp['is_training'] = False + else: + hp["is_training"] = False + + # Get checkpoint manager and load state + checkpoint_manager = components.get("checkpoint_manager") + if checkpoint_manager is None: + return pb2.RestoreCheckpointResponse( + success=False, + message="Checkpoint manager not initialized" + ) + + # Load checkpoint by hash + success = checkpoint_manager.load_state(experiment_hash) + + if success: + logger.info(f"Successfully restored checkpoint: {experiment_hash}") + return pb2.RestoreCheckpointResponse( + success=True, + message=f"Checkpoint {experiment_hash} restored successfully" + ) + else: + logger.warning(f"Failed to restore checkpoint: {experiment_hash}") + return pb2.RestoreCheckpointResponse( + success=False, + message=f"Failed to restore checkpoint {experiment_hash}" + ) + except Exception as e: + logger.error(f"Error during checkpoint restore: {str(e)}") + return pb2.RestoreCheckpointResponse( + success=False, + message=f"Error: {str(e)}" + ) diff --git a/weightslab/trainer/trainer_services.py b/weightslab/trainer/trainer_services.py index 98aa630f..ca7323a0 100644 --- a/weightslab/trainer/trainer_services.py +++ b/weightslab/trainer/trainer_services.py @@ -30,29 +30,9 @@ def __init__(self, exp_name: str = None, exp_service: ExperimentService | None = if exp_service is None: ctx = ExperimentContext(exp_name=exp_name) exp_service = ExperimentService(ctx=ctx) + self._ctx = ctx self._exp_service = exp_service - # ------------------------------------------------------------------------- - # Training status stream - # ------------------------------------------------------------------------- - def StreamStatus(self, request_iterator, context): - logger.debug(f"ExperimentServiceServicer.StreamStatus({request_iterator})") - - # # Get context components to fetch signal logger - # self._ctx.ensure_components() - # components = self._ctx.components - # is_model_interfaced = components.get("model") is not None - - # # delegate to domain ExperimentService - # if not is_model_interfaced: - # logger.warning("No signal_logger found in context components for StreamStatus") - # return None - - # # stream status updates to client - # for status in self._exp_service.stream_status(request_iterator): - # yield status - return self._exp_service.StreamStatus(request_iterator, context) - # ------------------------------------------------------------------------- # Sample retrieval (images / segmentation / recon) # ------------------------------------------------------------------------- @@ -97,6 +77,13 @@ def CheckAgentHealth(self, request, context): logger.debug(f"ExperimentServiceServicer.CheckAgentHealth({request})") return self._exp_service.data_service.CheckAgentHealth(request, context) + # ------------------------------------------------------------------------- + # Logger data sync for WeightsStudio + # ------------------------------------------------------------------------- + def GetLatestLoggerData(self, request, context): + logger.debug(f"ExperimentServiceServicer.GetLatestLoggerData({request})") + return self._exp_service.GetLatestLoggerData(request, context) + # ------------------------------------------------------------------------- # Training & hyperparameter commands # ------------------------------------------------------------------------- @@ -111,6 +98,13 @@ def ManipulateWeights(self, request, context): logger.debug(f"ExperimentServiceServicer.ManipulateWeights({request})") return self._exp_service.model_service.ManipulateWeights(request, context) + # ------------------------------------------------------------------------- + # Checkpoint restore + # ------------------------------------------------------------------------- + def RestoreCheckpoint(self, request, context): + logger.debug(f"ExperimentServiceServicer.RestoreCheckpoint({request})") + return self._exp_service.RestoreCheckpoint(request, context) + # ----------------------------------------------------------------------------- # Serving gRPC communication diff --git a/weightslab/trainer/trainer_tools.py b/weightslab/trainer/trainer_tools.py index 42404b00..14d62049 100644 --- a/weightslab/trainer/trainer_tools.py +++ b/weightslab/trainer/trainer_tools.py @@ -98,12 +98,12 @@ def get_neuron_representations(layer) -> Iterable[pb2.NeuronStatistics]: def get_layer_representation(layer) -> pb2.LayerRepresentation: layer_representation = None - layer_id = layer.get_module_id(), - layer_name = layer.__class__.__name__, - layer_type = layer.module_name, - incoming_neurons_count = layer.in_neurons, - neurons_count = layer.out_neurons, - kernel_size = (layer.kernel_size[0] if not isinstance(layer.kernel_size, (int, float)) else layer.kernel_size) if hasattr(layer, 'kernel_size') else None, + layer_id = layer.get_module_id() + layer_name = layer.__class__.__name__ + layer_type = layer.module_name + incoming_neurons_count = layer.in_neurons + neurons_count = layer.out_neurons + kernel_size = (layer.kernel_size[0] if not isinstance(layer.kernel_size, (int, float)) else layer.kernel_size) if hasattr(layer, 'kernel_size') else None stride = (layer.stride[0] if not isinstance(layer.stride, (int, float)) else layer.stride) if hasattr(layer, 'stride') else None parameters = { diff --git a/weightslab/utils/__init__.py b/weightslab/utils/__init__.py index f6052446..c1455d0d 100644 --- a/weightslab/utils/__init__.py +++ b/weightslab/utils/__init__.py @@ -1,3 +1,3 @@ -from .tools import filter_kwargs_for_callable, safe_call_with_kwargs +from .tools import filter_kwargs_for_callable, safe_call_with_kwargs, capture_rng_state, restore_rng_state -__all__ = ['filter_kwargs_for_callable', 'safe_call_with_kwargs'] +__all__ = ['filter_kwargs_for_callable', 'safe_call_with_kwargs', 'capture_rng_state', 'restore_rng_state'] diff --git a/weightslab/utils/board.py b/weightslab/utils/board.py deleted file mode 100644 index 1beef725..00000000 --- a/weightslab/utils/board.py +++ /dev/null @@ -1,31 +0,0 @@ -import queue - - -class Dash: - def __init__(self) -> None: - self.queue = queue.Queue() - self.graph_names = set() - - def get_graph_names(self): - return list(self.graph_names) - - def add_scalars(self, graph_name, name_2_value, global_step: int): - self.graph_names.add(graph_name) - for line_name, line_value in name_2_value.items(): - self.queue.put({ - "experiment_name": line_name, - "model_age": global_step, - "metric_name": graph_name, - "metric_value": float(line_value), - }) - - def add_annotations( - self, graph_names, line_name, annotation, global_step, - metadata=None): - - self.queue.put({ - "experiment_name": line_name, - "model_age": global_step, - "annotation": annotation, - "metadata": metadata, - }) diff --git a/weightslab/utils/computational_graph.py b/weightslab/utils/computational_graph.py index ce08f45a..c7f8c5eb 100644 --- a/weightslab/utils/computational_graph.py +++ b/weightslab/utils/computational_graph.py @@ -709,26 +709,26 @@ def generate_graph_dependencies_from_torchfx( # directly instead of passing through to an earlier source. if current_module is not None: node_to_module[node] = make_safelist(current_module) - - # SEED NEURONS: Use FX metadata to seed neurons if possible - for mod in make_safelist(current_module): - if 'tensor_meta' in node.meta: - meta = node.meta['tensor_meta'] - if hasattr(meta, 'shape') and len(meta.shape) >= 2: - out_ch = meta.shape[1] - if out_ch is not None and out_ch > 0: - mod.set_neurons('out_neurons', out_ch) - if getattr(mod, 'wl_same_flag', False): - mod.set_neurons('in_neurons', out_ch) - - # Also check inputs to seed in_neurons - for arg in node.args: - if isinstance(arg, th.fx.Node) and 'tensor_meta' in arg.meta: - meta_in = arg.meta['tensor_meta'] - if hasattr(meta_in, 'shape') and len(meta_in.shape) >= 2: - in_ch = meta_in.shape[1] - if in_ch is not None and in_ch > 0: - mod.set_neurons('in_neurons', in_ch) + + # # SEED NEURONS: Use FX metadata to seed neurons if possible + # for mod in make_safelist(current_module): + # if 'tensor_meta' in node.meta: + # meta = node.meta['tensor_meta'] + # if hasattr(meta, 'shape') and len(meta.shape) >= 2: + # out_ch = meta.shape[1] + # if out_ch is not None and out_ch > 0: + # mod.set_neurons('out_neurons', out_ch) + # if getattr(mod, 'wl_same_flag', False): + # mod.set_neurons('in_neurons', out_ch) + + # # Also check inputs to seed in_neurons + # for arg in node.args: + # if isinstance(arg, th.fx.Node) and 'tensor_meta' in arg.meta: + # meta_in = arg.meta['tensor_meta'] + # if hasattr(meta_in, 'shape') and len(meta_in.shape) >= 2: + # in_ch = meta_in.shape[1] + # if in_ch is not None and in_ch > 0: + # mod.set_neurons('in_neurons', in_ch) # --- Handle General Merge Operations (Any call_function with multiple # module inputs) --- @@ -1193,16 +1193,16 @@ def module_for_tensor(tname: str) -> Optional[nn.Module]: src_channels = get_channel_count(src_tensor) dst_channels = get_channel_count(dst_tensor) if dst_tensor else None - # SEED NEURONS: Use ONNX metadata to seed neurons if possible - if src_channels is not None and src_channels > 0: - src_mod.set_neurons('out_neurons', src_channels) - if getattr(src_mod, 'wl_same_flag', False): - src_mod.set_neurons('in_neurons', src_channels) - - if dst_channels is not None and dst_channels > 0: - dst_mod.set_neurons('in_neurons', dst_channels) - if getattr(dst_mod, 'wl_same_flag', False): - dst_mod.set_neurons('out_neurons', dst_channels) + # # SEED NEURONS: Use ONNX metadata to seed neurons if possible + # if src_channels is not None and src_channels > 0: + # src_mod.set_neurons('out_neurons', src_channels) + # if getattr(src_mod, 'wl_same_flag', False): + # src_mod.set_neurons('in_neurons', src_channels) + + # if dst_channels is not None and dst_channels > 0: + # dst_mod.set_neurons('in_neurons', dst_channels) + # if getattr(dst_mod, 'wl_same_flag', False): + # dst_mod.set_neurons('out_neurons', dst_channels) logger.debug(f"Analyzing dependency {src_name} -> {dst_name}") logger.debug(f" Source channels: {src_channels}, Destination channels: {dst_channels}") diff --git a/weightslab/utils/logger.py b/weightslab/utils/logger.py new file mode 100644 index 00000000..476b3fb8 --- /dev/null +++ b/weightslab/utils/logger.py @@ -0,0 +1,106 @@ +import queue + +from weightslab.backend.ledgers import get_logger, register_logger, get_checkpoint_manager + + +class LoggerQueue: + def __init__(self, name: str = None, register: bool = True) -> None: + self.graph_names = set() + self._current_step_buffer = {} # {metric_name: [values]} + self._last_step = None + self._signal_history = [] # Keep all signals in memory for persistence + self._pending_queue = [] # Queue for new signals waiting to be sent to WeightsStudio + + if register: + try: + get_logger(name) + except Exception: + pass + register_logger(name, self) + + self.chkpt_manager = get_checkpoint_manager() + + def get_graph_names(self): + return list(self.graph_names) + + def _flush_step_buffer(self): + """Flush accumulated metrics for the previous step to history.""" + if self._current_step_buffer and self._last_step is not None: + for metric_name, values in self._current_step_buffer.items(): + signal = { + "experiment_name": metric_name, + "model_age": self._last_step, + "metric_name": metric_name, + "metric_value": sum(values) / len(values) if len(values) > 1 else values[0], + "experiment_hash": self.chkpt_manager.get_current_experiment_hash() if self.chkpt_manager else None, + } + self._signal_history.append(signal) + # Also add to pending queue for WeightsStudio + self._pending_queue.append(signal) + self._current_step_buffer.clear() + + def add_scalars(self, graph_name, name_2_value, global_step: int): + global_step -= 1 # adjust for 0-based step indexing + self.graph_names.add(graph_name) + + # If step changed, flush the previous step's buffer + if global_step != self._last_step: + self._flush_step_buffer() + self._last_step = global_step + + for _, line_value in name_2_value.items(): + metric_key = f"{graph_name}" + if metric_key not in self._current_step_buffer: + self._current_step_buffer[metric_key] = [] + self._current_step_buffer[metric_key].append(float(line_value)) + + def print_history(self): + """Print all items in history.""" + for i, item in enumerate(self._signal_history): + print(f"[{i}] {item}") + return self._signal_history + + def print_buffer(self): + """Print current step buffer contents.""" + print(f"Current step: {self._last_step}") + print(f"Buffered metrics: {self._current_step_buffer}") + return self._current_step_buffer + + def get_signal_history(self): + """Retrieve all accumulated signals from memory.""" + return list(self._signal_history) + + def get_and_clear_queue(self): + """Get pending queue and clear it (for incremental updates to WeightsStudio).""" + queue_copy = list(self._pending_queue) + self._pending_queue.clear() + return queue_copy + + def load_signal_history(self, signals): + """Load a list of signals into history (used for checkpoint restore).""" + if not signals: + return + for signal in signals: + self._signal_history.append(signal) + try: + metric_name = signal.get("metric_name") + if metric_name: + # Derive a graph name if encoded as 'graph:metric' + if ":" in metric_name: + graph, _ = metric_name.split(":", 1) + self.graph_names.add(graph) + except Exception: + continue + + def load_snapshot(self, snapshot: dict): + """Restore logger state from a snapshot dict.""" + if not snapshot: + return + graph_names = snapshot.get("graph_names", []) + self.graph_names.update(graph_names) + signals = snapshot.get("signal_history", []) + self.load_signal_history(signals) + + def clear_signal_history(self): + """Clear signal history.""" + self._signal_history.clear() diff --git a/weightslab/utils/tools.py b/weightslab/utils/tools.py index fed06517..530758dd 100644 --- a/weightslab/utils/tools.py +++ b/weightslab/utils/tools.py @@ -32,13 +32,112 @@ def release(self): # ---------------------------------------------------------------------------- # -------------------------- Utils Functions --------------------------------- # ---------------------------------------------------------------------------- -def seed_everything(seed): +def seed_everything(seed=42): """Seed everything for reproducibility.""" np.random.seed(seed) th.manual_seed(seed) random.seed(seed) th.backends.cudnn.deterministic = True + # Reproducibility + rng = capture_rng_state() + print(rng) + restore_rng_state(rng) + + +def capture_rng_state(): + """ + Capture all RNG states in a JSON-serializable format. + + Returns: + dict: Dictionary with python_random, numpy_random, torch_rng, torch_cuda_rng + in serializable formats (lists/tuples) + """ + rng_state = { + 'python_random': random.getstate(), + 'torch_rng': th.get_rng_state().cpu().tolist() if hasattr(th.get_rng_state(), 'tolist') else str(th.get_rng_state()), + } + + # NumPy random state: (version, internal_state_array, gauss_next) + np_state = np.random.get_state() + rng_state['numpy_random'] = (np_state[0], np_state[1].tolist(), np_state[2]) + + # Add CUDA RNG if available + if th.cuda.is_available(): + try: + rng_state['torch_cuda_rng'] = th.cuda.get_rng_state().cpu().tolist() if hasattr(th.cuda.get_rng_state(), 'tolist') else str(th.cuda.get_rng_state()) + except Exception as e: + logger.warning(f"Failed to capture CUDA RNG state: {e}") + + return rng_state + +def restore_rng_state(rng_state): + """ + Restore RNG states from captured state dictionary. + + Args: + rng_state (dict): Dictionary with RNG states (python_random, numpy_random, torch_rng, torch_cuda_rng) + """ + if not rng_state: + logger.warning("RNG state is None or empty, skipping restoration") + return + + try: + # Restore Python random state + if 'python_random' in rng_state: + try: + random.setstate(tuple(tuple(i) if i is not None and not isinstance(i, (int, float)) else i for i in rng_state['python_random'])) # Conver to tuple of tuples + logger.debug("Restored Python random state") + except Exception as e: + logger.warning(f"Failed to restore Python random state: {e}") + + # Restore NumPy random state + if 'numpy_random' in rng_state: + try: + state_data = rng_state['numpy_random'] + if isinstance(state_data, (list, tuple)) and len(state_data) == 3: + version, internal, gauss = state_data + if isinstance(internal, list): + internal = np.array(internal, dtype=np.uint32) + np.random.set_state((version, internal, gauss)) + logger.debug("Restored NumPy random state") + else: + logger.warning("NumPy RNG state format invalid") + except Exception as e: + logger.warning(f"Failed to restore NumPy random state: {e}") + + # Restore PyTorch RNG state + if 'torch_rng' in rng_state: + try: + torch_state = rng_state['torch_rng'] + if isinstance(torch_state, list): + torch_state = th.tensor(torch_state, dtype=th.uint8) + elif isinstance(torch_state, str): + logger.warning("Torch RNG state is a string representation, cannot restore") + return + th.set_rng_state(torch_state) + logger.debug("Restored PyTorch RNG state") + except Exception as e: + logger.warning(f"Failed to restore PyTorch RNG state: {e}") + + # Restore CUDA RNG state if available + if 'torch_cuda_rng' in rng_state and th.cuda.is_available(): + try: + cuda_state = rng_state['torch_cuda_rng'] + if isinstance(cuda_state, list): + cuda_state = th.tensor(cuda_state, dtype=th.uint8) + elif isinstance(cuda_state, str): + logger.debug("CUDA RNG state is a string representation, skipping") + return + th.cuda.set_rng_state(cuda_state) + logger.debug("Restored CUDA RNG state") + except Exception as e: + logger.warning(f"Failed to restore CUDA RNG state: {e}") + + logger.debug("Successfully restored RNG states") + except Exception as e: + logger.error(f"Error restoring RNG state: {e}") + def extract_in_out_params(module: nn.Module) -> List[int | str]: """ Detects and returns the primary input and output dimension parameters @@ -80,7 +179,7 @@ def extract_in_out_params(module: nn.Module) -> List[int | str]: # 4. Pass-through layers (Pooling, Upsampling, Dropout, Activations) # These layers maintain the same number of channels/neurons. pass_through_types = [ - 'Pool', 'Upsample', 'Dropout', 'ReLU', 'PReLU', 'LeakyReLU', + 'Pool', 'Upsample', 'Dropout', 'ReLU', 'PReLU', 'LeakyReLU', 'Sigmoid', 'Tanh', 'ELU', 'Softmax', 'Identity' ] module_name = module._get_name()