| name | gere-continual-learning |
| title | GeRe - Anti-Forgetting in Continual Learning of LLMs |
| version | 0.0.2 |
| engine | skillxiv-v0.0.2-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2508.04676 |
| keywords | ["continual-learning","catastrophic-forgetting","replay-memory","llm-training"] |
| description | Prevents catastrophic forgetting in continual LLM learning using threshold-based margin loss with fixed general replay samples from pretraining data. |
GeRe: Anti-Forgetting in Continual Learning of LLMs
Core Concept
GeRe addresses catastrophic forgetting in continual LLM learning by employing a threshold-based margin (TM) loss function on fixed general replay samples derived from pretraining. This maintains activation state consistency during replay learning, effectively mitigating both forgetting of general capabilities and performance degradation on previously learned tasks.
Architecture Overview
- General Sample Replay: Small fixed set of pretraining samples
- Threshold-Based Margin Loss: Constrains activation states during replay
- Task-Incremental Learning: Learn new domains sequentially
- Activation State Consistency: Maintain neural patterns from pretraining
Implementation Steps
Step 1: Collect General Replay Samples
Curate pretraining samples for replay:
class GeneralSampleCollector:
def __init__(self, pretraining_data, sample_size=1000):
super().__init__()
self.sample_size = sample_size
def select_general_samples(self, pretraining_data):
"""Select representative pretraining samples."""
samples = []
topics = self._identify_topics(pretraining_data)
for topic in topics:
topic_samples = [d for d in pretraining_data if d['topic'] == topic]
num_select = self.sample_size // len(topics)
selected = random.sample(topic_samples, min(num_select, len(topic_samples)))
samples.extend(selected)
return samples[:self.sample_size]
def _identify_topics(self, data):
"""Identify topic distribution."""
topics = {}
for sample in data[:100]:
topic = sample.get('topic', 'general')
topics[topic] = topics.get(topic, 0) + 1
return topics.keys()
Step 2: Implement Threshold-Based Margin Loss
Design loss for activation consistency:
class ThresholdBasedMarginLoss(nn.Module):
def __init__(self, threshold=0.5, margin=0.1):
super().__init__()
self.threshold = threshold
self.margin = margin
def compute_loss(self, original_activations, replay_activations):
"""
Compute TM loss maintaining activation state consistency.
Args:
original_activations: Activations from original training
replay_activations: Activations during replay
Returns:
loss: TM loss value
"""
orig_norm = F.normalize(original_activations, dim=-1)
replay_norm = F.normalize(replay_activations, dim=-1)
similarity = torch.mm(replay_norm, orig_norm.t())
loss = 0
for i in range(similarity.shape[0]):
for j in range(similarity.shape[1]):
sim = similarity[i, j]
if sim > self.threshold:
loss += torch.relu(self.margin - sim)
else:
loss += torch.relu(sim - self.threshold)
return loss / (similarity.shape[0] * similarity.shape[])
Step 3: Implement Continual Learning Loop
Train on sequential tasks:
class ContinualLearningTrainer:
def __init__(self, model, general_samples):
super().__init__()
self.model = model
self.general_samples = general_samples
self.activation_history = {}
def learn_new_task(self, new_task_data, task_id):
"""Learn new task while replaying general samples."""
optimizer = AdamW(self.model.parameters(), lr=2e-5)
tm_loss_fn = ThresholdBasedMarginLoss()
for epoch in range(3):
batch_new = random.sample(new_task_data, min(32, len(new_task_data)))
batch_replay = random.sample(self.general_samples, min(32, len(self.general_samples)))
for example in batch_new:
outputs = self.model(example['input_ids'], output_hidden_states=True)
task_loss = F.cross_entropy(outputs.logits, example['labels'])
activations = outputs.hidden_states[-1]
if task_id == 0:
self.activation_history[example[]] = activations.detach()
optimizer.zero_grad()
task_loss.backward()
optimizer.step()
example batch_replay:
outputs = .model(example[], output_hidden_states=)
current_activations = outputs.hidden_states[-]
example[] .activation_history:
original = .activation_history[example[]]
:
original = current_activations.detach()
tm_loss = tm_loss_fn.compute_loss(original, current_activations)
optimizer.zero_grad()
tm_loss.backward()
optimizer.step()
():
results = {}
task_id, task_data task_data_dict.items():
correct =
total = (task_data)
torch.no_grad():
example task_data:
outputs = .model(example[])
pred = outputs.logits.argmax(dim=-)
pred == example[]:
correct +=
results[] = correct / total
results
Practical Guidance
Hyperparameters and Configuration:
- General sample size: 1000-5000 samples
- TM threshold: 0.5
- TM margin: 0.1-0.2
- Learning rate: 2e-5 to 5e-5
- Replay batch size: 32
When to Use GeRe:
- Continual learning of LLMs on sequential domains
- Systems where forgetting general knowledge is costly
- Scenarios with limited task-specific data
- Multi-domain adaptation requiring stability
When NOT to Use:
- Single-domain training (no continual aspect)
- When task-specific performance is prioritized over general knowledge
- Very large models (storage overhead)
Implementation Notes:
- Fixed replay samples prevent distribution shifts
- TM loss more robust than standard replay approaches
- Monitor forgetting vs learning tradeoff
- Consider optimal buffer size vs memory constraints
Reference
Paper: GeRe: Anti-Forgetting in Continual Learning of LLM
ArXiv: 2508.04676
Performance: Consistently improves upon label fitting and KL divergence replay baselines