A recurrent state belongs to one ordered stream; detaching it limits gradient history, while resetting it prevents another stream from inheriting private context.
Recurrent hidden-state resets and truncated gradients
Define state ownership
An LSTM carries hidden and cell tensors with layer, batch and width axes. Batch column zero at the next chunk must still represent the same stream as column zero in the previous chunk, or the state must be reset. Shuffling chunks without their stream identifiers creates cross-customer contamination even when the loss calculation succeeds. At serving time, key state by an opaque conversation or device identifier, put a time-to-live on it, and clear it at an end event or identity change. A state cache is part of the prediction contract, not an internal convenience.
Detach at the training boundary
If chunks from one long history are processed in order, the computational graph can span all previous chunks unless hidden and cell states are detached after a chosen truncation boundary. Detaching keeps their numeric contents for the next chunk while stopping gradients from flowing backward across that boundary. Resetting to zero discards the contents too. These operations solve different problems. Gradient accumulation can still combine losses from multiple chunks before an optimizer step, but retained graphs and update timing must be accounted for.
Apply per-stream reset masks
A batch can contain one continuing device and one new device. Reset only the changed column, not the whole batch, by multiplying detached state by a zero-or-one keep mask broadcast over layers and features. Verify the mask comes from stable stream identity, not from a label or sequence length that merely happens to match. If a stream is removed and the batch is compacted, reorder state columns with the same mapping or allocate a fresh state. A checksum of stream IDs next to state helps expose accidental reuse.
Name the temporal limit
Truncating every 23 events means a loss cannot send gradient through earlier chunks, although the numeric state can carry old information forward. That is a modeling trade: shorter spans reduce memory and make updates more frequent, but may weaken learning of long dependencies. Compare several spans against a fixed validation set grouped by stream and report event count, optimizer updates and peak memory. Packing handles within-batch padding; it does not decide how far the backward graph reaches.
Test a handoff failure
Feed stream A for one chunk, replace its identifier with stream B and assert B receives a zero initial state. Feed a second chunk of A and assert the state is retained but detached. Hold model weights fixed in evaluation mode when testing state ownership; otherwise dropout or training updates can make the comparison noisy. Also test worker restart: either restore state from a versioned checkpoint or explicitly start a new session, with a documented effect on predictions.
Implementation
import torch
from torch import nn
torch.manual_seed(83)
event_lstm = nn.LSTM(input_size=4, hidden_size=6, batch_first=True)
event_head = nn.Linear(6, 2)
optimizer = torch.optim.AdamW(list(event_lstm.parameters())
+ list(event_head.parameters()), lr=0.00047)
first_chunk = torch.randn(2, 3, 4)
first_labels = torch.tensor([0, 1])
optimizer.zero_grad(set_to_none=True)
first_outputs, carried_state = event_lstm(first_chunk)
first_loss = nn.functional.cross_entropy(event_head(first_outputs[:, -1]),
first_labels)
first_loss.backward()
optimizer.step()
same_stream_next = torch.tensor([True, False]).view(1, 2, 1)
next_state = tuple(state.detach() * same_stream_next for state in carried_state)
assert all(not state.requires_grad for state in next_state)
assert torch.count_nonzero(next_state[0][:, 1]) == 0
second_chunk = torch.randn(2, 3, 4)
second_outputs, second_state = event_lstm(second_chunk, next_state)
assert second_outputs.shape == (2, 3, 6)
assert second_state[0].shape == (1, 2, 6)Performance and operating cost
Without truncation, retained activations and graph history grow with the number of processed steps until backward releases them. Detaching bounds that history to the chosen chunk span, while state tensors themselves occupy O(layers × batch × hidden width). LSTM work remains sequential across timesteps. An online state cache adds memory proportional to concurrent streams, so expiration and eviction belong in capacity planning.
Common Mistakes
- Do not carry a hidden state into a different stream after batch reshuffling.
- Do not confuse detach, which cuts gradients, with reset, which clears information.
- Do not claim long-horizon training when the backward span ends at every short chunk.
Read next
- Packed LSTM sequences and the last valid state
- Gradient accumulation and effective batch accounting
- Checkpoint recovery: save optimizer state and the run boundary
- Inference contracts: preserve preprocessing and measure tail latency
- Project: predict service escalation from event histories
Continue the workflow: Causal dilated convolution and receptive field.
