| name | inputs-and-layer-wise-states |
| description | Analyze and visualize layer-wise gradient behaviors in LLMs during fine-tuning for fast vs slow thinking tasks, calculate gradient statistics, and understand training patterns across different model layers |
Demo Scripts
scripts/calculate_gradients.py
"""
Calculate Layer-wise Gradient Statistics for LLM Fine-tuning
This script demonstrates how to calculate gradient statistics for each layer
when fine-tuning LLMs on different types of responses (fast vs slow thinking).
"""
import json
import torch
import numpy as np
from typing import Dict, List, Tuple, Optional
from transformers import AutoTokenizer, AutoModelForCausalLM
import argparse
from pathlib import Path
def load_training_data(data_path: str) -> List[Dict]:
"""
Load training data from JSON file.
Args:
data_path: Path to the JSON data file
Returns:
List of training examples
"""
with open(data_path, 'r') as f:
data = json.load(f)
return data
def prepare_model_and_tokenizer(
model_name_or_path: str,
device: str = "cuda"
) -> Tuple[AutoModelForCausalLM, AutoTokenizer]:
"""
Load and prepare model and tokenizer for gradient calculation.
Args:
model_name_or_path: Hugging Face model identifier or local path
device: Device to load model on
Returns:
Tuple of (model, tokenizer)
"""
tokenizer = AutoTokenizer.from_pretrained(model_name_or_path)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
model = AutoModelForCausalLM.from_pretrained(
model_name_or_path,
torch_dtype=torch.float16,
device_map="auto"
)
model.eval()
return model, tokenizer
def calculate_layer_gradients(
model: AutoModelForCausalLM,
tokenizer: AutoTokenizer,
text: str,
max_length: int = 1024
) -> Dict[str, float]:
"""
Calculate gradient norms for each layer of the model.
Args:
model: The language model
tokenizer: The tokenizer
text: Input text for gradient calculation
max_length: Maximum sequence length
Returns:
Dictionary mapping layer names to gradient norms
"""
inputs = tokenizer(
text,
return_tensors="pt",
truncation=True,
max_length=max_length,
padding=True
).to(model.device)
model.zero_grad()
with torch.enable_grad():
outputs = model(**inputs, labels=inputs["input_ids"])
loss = outputs.loss
loss.backward()
gradient_norms = {}
for name, param in model.named_parameters():
if param.grad is not None:
grad_norm = torch.norm(param.grad, p=2).item()
gradient_norms[name] = grad_norm
return gradient_norms
def calculate_svd_vectors(
model: AutoModelForCausalLM,
tokenizer: AutoTokenizer,
text: str,
num_components: int = 10
) -> Dict[str, np.ndarray]:
"""
Calculate SVD vectors for gradient analysis.
Args:
model: The language model
tokenizer: The tokenizer
text: Input text
num_components: Number of SVD components to compute
Returns:
Dictionary mapping layer names to SVD components
"""
inputs = tokenizer(text, return_tensors="pt", truncation=True).to(model.device)
model.zero_grad()
with torch.enable_grad():
outputs = model(**inputs, labels=inputs["input_ids"])
outputs.loss.backward()
svd_results = {}
for name, param in model.named_parameters():
if param.grad is not None and len(param.grad.shape) >= 2:
grad_flat = param.grad.view(param.grad.shape[0], -1).cpu().numpy()
try:
U, S, Vt = np.linalg.svd(grad_flat, full_matrices=False)
svd_results[name] = {
'singular_values': S[:num_components].tolist(),
'top_component_variance': (S[0]**2 / np.sum(S**2)).item()
}
except:
svd_results[name] = None
return svd_results
def analyze_gradient_patterns(
gradient_norms: Dict[str, float],
layer_groups: Optional[Dict[str, List[str]]] = None
) -> Dict[str, float]:
"""
Analyze gradient patterns across layers.
Args:
gradient_norms: Dictionary of layer gradient norms
layer_groups: Optional grouping of layers (e.g., early, middle, late)
Returns:
Dictionary of gradient statistics
"""
norms = list(gradient_norms.values())
stats = {
'mean_norm': np.mean(norms),
'std_norm': np.std(norms),
'max_norm': np.max(norms),
'min_norm': np.min(norms),
'coefficient_of_variation': np.std(norms) / np.mean(norms) if np.mean(norms) > 0 else 0
}
if len(norms) > 1:
differences = [abs(norms[i+1] - norms[i]) for i in range(len(norms)-1)]
stats['mean_layer_difference'] = np.mean(differences)
stats['max_layer_difference'] = np.max(differences)
if layer_groups:
for group_name, layer_names in layer_groups.items():
group_norms = [gradient_norms[name] for name in layer_names if name in gradient_norms]
if group_norms:
stats[f'{group_name}_mean'] = np.mean(group_norms)
stats[f'{group_name}_std'] = np.std(group_norms)
return stats
def process_dataset(
model: AutoModelForCausalLM,
tokenizer: AutoTokenizer,
data: List[Dict],
output_path: str,
max_samples: int = None
):
"""
Process entire dataset and save gradient statistics.
Args:
model: The language model
tokenizer: The tokenizer
data: List of training examples
output_path: Path to save results
max_samples: Maximum number of samples to process
"""
results = []
if max_samples:
data = data[:max_samples]
for idx, example in enumerate(data):
print(f"Processing example {idx+1}/{len(data)}")
if 'instruction' in example and 'response' in example:
text = f"{example['instruction']}\n{example['response']}"
elif 'text' in example:
text = example['text']
else:
continue
gradient_norms = calculate_layer_gradients(model, tokenizer, text)
stats = analyze_gradient_patterns(gradient_norms)
result = {
'example_id': idx,
'gradient_norms': gradient_norms,
'statistics': stats
}
results.append(result)
if (idx + 1) % 10 == 0:
with open(output_path, 'w') as f:
for res in results:
f.write(json.dumps(res) + '\n')
with open(output_path, 'w') as f:
for res in results:
f.write(json.dumps(res) + '\n')
print(f"Results saved to {output_path}")
def main():
parser = argparse.ArgumentParser(description="Calculate layer-wise gradient statistics")
parser.add_argument("--data_path", type=str, required=True, help="Path to training data")
parser.add_argument("--model_name_or_path", type=str, required=True, help="Model identifier")
parser.add_argument("--output_path", type=str, required=True, help="Output path for results")
parser.add_argument("--max_samples", type=int, default=None, help="Maximum samples to process")
parser.add_argument("--max_length", type=int, default=1024, help="Maximum sequence length")
args = parser.parse_args()
print(f"Loading data from {args.data_path}")
data = load_training_data(args.data_path)
print(f"Loading model: {args.model_name_or_path}")
model, tokenizer = prepare_model_and_tokenizer(args.model_name_or_path)
process_dataset(
model=model,
tokenizer=tokenizer,
data=data,
output_path=args.output_path,
max_samples=args.max_samples
)
if __name__ == "__main__":
main()
scripts/visualize_gradients.py
"""
Visualize Layer-wise Gradient Statistics
This script provides visualization capabilities for gradient statistics
calculated during LLM fine-tuning, comparing fast vs slow thinking patterns.
"""
import json
import numpy as np
import matplotlib.pyplot as plt
from typing import Dict, List, Optional, Tuple
import seaborn as sns
from pathlib import Path
import pandas as pd
def load_gradient_results(jsonl_path: str) -> List[Dict]:
"""
Load gradient results from JSONL file.
Args:
jsonl_path: Path to JSONL file containing gradient statistics
Returns:
List of gradient result dictionaries
"""
results = []
with open(jsonl_path, 'r') as f:
for line in f:
results.append(json.loads(line.strip()))
return results
def extract_layer_gradients(results: List[Dict]) -> pd.DataFrame:
"""
Extract and organize layer gradients into a DataFrame.
Args:
results: List of gradient results
Returns:
DataFrame with layer gradients
"""
data = []
for result in results:
gradient_norms = result.get('gradient_norms', {})
layer_name, norm gradient_norms.items():
layer_num = extract_layer_number(layer_name)
data.append({
: result.get(, ),
: layer_name,
: layer_num,
: norm
})
pd.DataFrame(data)
() -> :
re
= re.search(, layer_name)
:
(.group())
() -> [, ]:
mad_stats = {}
section_size = (values) // num_sections
i (num_sections):
start = i * section_size
end = (i + ) * section_size i < num_sections - (values)
section = values[start:end]
(section) > :
mean = np.mean(section)
mad = np.mean([(x - mean) x section])
mad_stats[] = mad
mad_stats[] = mean
mad_stats
():
fig, axes = plt.subplots(, , figsize=(, ))
ax = axes[]
fast_mean = df_fast.groupby()[].mean()
slow_mean = df_slow.groupby()[].mean()
ax.plot(fast_mean.index, fast_mean.values, label=, marker=, linewidth=)
ax.plot(slow_mean.index, slow_mean.values, label=, marker=, linewidth=)
ax.set_xlabel(, fontsize=)
ax.set_ylabel(, fontsize=)
ax.set_title(, fontsize=)
ax.legend()
ax.grid(, alpha=)
ax = axes[]
fast_std = df_fast.groupby()[].std()
slow_std = df_slow.groupby()[].std()
ax.bar(fast_std.index - , fast_std.values, width=, label=, alpha=)
ax.bar(slow_std.index + , slow_std.values, width=, label=, alpha=)
ax.set_xlabel(, fontsize=)
ax.set_ylabel(, fontsize=)
ax.set_title(, fontsize=)
ax.legend()
ax.grid(, alpha=)
plt.tight_layout()
save_path:
plt.savefig(save_path, dpi=, bbox_inches=)
()
plt.show()
():
pivot_data = df.pivot_table(
index=,
columns=,
values=,
aggfunc=
)
plt.figure(figsize=(, ))
sns.heatmap(
pivot_data,
cmap=,
cbar_kws={: },
xticklabels=,
yticklabels=
)
plt.xlabel(, fontsize=)
plt.ylabel(, fontsize=)
plt.title(title, fontsize=)
save_path:
plt.savefig(save_path, dpi=, bbox_inches=)
()
plt.show()
() -> []:
differences = []
v1, v2 (values1, values2):
v2 != :
diff = (v1 - v2) / (v2)
:
diff = (v1) v1 !=
differences.append(diff)
differences
() -> pd.DataFrame:
fast_stats = aggregate_statistics(fast_results)
slow_stats = aggregate_statistics(slow_results)
comparison = {
: [],
: [],
: [],
: []
}
metric fast_stats.keys():
comparison[].append(metric)
comparison[].append(fast_stats[metric])
comparison[].append(slow_stats[metric])
slow_stats[metric] != :
rel_diff = (fast_stats[metric] - slow_stats[metric]) / (slow_stats[metric])
:
rel_diff =
comparison[].append(rel_diff)
pd.DataFrame(comparison)
() -> [, ]:
all_stats = {}
result results:
stats = result.get(, {})
key, value stats.items():
key all_stats:
all_stats[key] = []
all_stats[key].append(value)
aggregated = {}
key, values all_stats.items():
aggregated[key] = np.mean(values)
aggregated
():
max_layer = df[].()
section_size = (max_layer + ) // num_sections
df[] = df[].apply(
x: (x // section_size, num_sections - )
)
section_stats = df.groupby()[].agg([, , , ])
fig, axes = plt.subplots(, , figsize=(, ))
ax = axes[, ]
ax.bar(section_stats.index, section_stats[])
ax.set_xlabel()
ax.set_ylabel()
ax.set_title()
ax.set_xticks((num_sections))
ax.set_xticklabels([, , ][:num_sections])
ax = axes[, ]
ax.bar(section_stats.index, section_stats[], color=)
ax.set_xlabel()
ax.set_ylabel()
ax.set_title()
ax.set_xticks((num_sections))
ax.set_xticklabels([, , ][:num_sections])
ax = axes[, ]
df.boxplot(column=, by=, ax=ax)
ax.set_xlabel()
ax.set_ylabel()
ax.set_title()
ax.set_xticklabels([, , ][:num_sections])
plt.sca(ax)
plt.xticks((, num_sections + ), [, , ][:num_sections])
ax = axes[, ]
positions = (df[].unique())
parts = ax.violinplot(
[df[df[] == s][].values s positions],
positions=positions,
showmeans=,
showmedians=
)
ax.set_xlabel()
ax.set_ylabel()
ax.set_title()
ax.set_xticks((num_sections))
ax.set_xticklabels([, , ][:num_sections])
plt.tight_layout()
save_path:
plt.savefig(save_path, dpi=, bbox_inches=)
()
plt.show()
():
argparse
parser = argparse.ArgumentParser(description=)
parser.add_argument(, =, required=,
=)
parser.add_argument(, =, required=,
=)
parser.add_argument(, =, default=,
=)
args = parser.parse_args()
output_dir = Path(args.output_dir)
output_dir.mkdir(parents=, exist_ok=)
()
fast_results = load_gradient_results(args.fast_gradients)
slow_results = load_gradient_results(args.slow_gradients)
df_fast = extract_layer_gradients(fast_results)
df_slow = extract_layer_gradients(slow_results)
()
plot_layer_gradient_comparison(
df_fast, df_slow,
save_path=output_dir /
)
()
plot_gradient_heatmap(
df_fast,
title=,
save_path=output_dir /
)
plot_gradient_heatmap(
df_slow,
title=,
save_path=output_dir /
)
()
visualize_layer_sections(
df_fast,
save_path=output_dir /
)
visualize_layer_sections(
df_slow,
save_path=output_dir /
)
()
comparison_table = generate_comparison_table(fast_results, slow_results)
comparison_table.to_csv(output_dir / , index=)
(comparison_table)
()
__name__ == :
main()