Skip to content

Commit e293a8b

Browse files
henryclaude
andcommitted
fix: size-adaptive memory pool to prevent OOM on 224x224 training
MemoryPool was designed for small CIFAR-10 buffers but used the same aggressive caching (refill=32, max=64 per class) for 224x224 tensors, causing multi-GB allocations. Three fixes: - Size-adaptive limits: buffers >=1MB get refill=2, max=4 (was 32/64) - trim() now flushes thread-local caches (previously only global pool) - Training loop scopes batch tensors so trim() can actually free them Also adds Tensor::live_count()/live_bytes() diagnostics and MemoryPool::cached_bytes() for tracking memory across batches. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
1 parent e2e5694 commit e293a8b

5 files changed

Lines changed: 138 additions & 53 deletions

File tree

‎core/memory_pool.cpp‎

Lines changed: 51 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -28,8 +28,20 @@ void MemoryPool::set_allocator(AllocFn alloc, FreeFn free) {
2828

2929
namespace {
3030

31-
const size_t kMaxPerClassPerThread = 64;
32-
const size_t kRefillBatch = 32;
31+
// Size-adaptive limits: large buffers (e.g. 128MB for [32,64,112,112])
32+
// should NOT fill thread-local caches with dozens of copies.
33+
// Threshold: 256K floats = 1MB. Above that, cache fewer buffers.
34+
const size_t kLargeBufferThreshold = 256 * 1024; // floats
35+
36+
static size_t max_per_class(size_t cls) {
37+
if (cls >= kLargeBufferThreshold) return 4;
38+
return 64;
39+
}
40+
41+
static size_t refill_batch(size_t cls) {
42+
if (cls >= kLargeBufferThreshold) return 2;
43+
return 32;
44+
}
3345

3446
size_t next_power_of_two(size_t n) {
3547
if (n == 0) return 1;
@@ -138,7 +150,7 @@ struct ThreadCache {
138150
return ptr;
139151
}
140152
std::vector<float*> batch;
141-
global->acquire_batch(cls, kRefillBatch, batch);
153+
global->acquire_batch(cls, refill_batch(cls), batch);
142154
if (batch.empty()) return nullptr;
143155
for (size_t i = 1; i < batch.size(); ++i)
144156
local.push_back(batch[i]);
@@ -150,11 +162,12 @@ struct ThreadCache {
150162
global_ = global;
151163
size_t cls = MemoryPool::size_class(original_n);
152164
auto& local = buckets[cls];
153-
if (local.size() < kMaxPerClassPerThread) {
165+
size_t max_local = max_per_class(cls);
166+
if (local.size() < max_local) {
154167
local.push_back(ptr);
155168
return;
156169
}
157-
size_t drain_count = kMaxPerClassPerThread / 2;
170+
size_t drain_count = max_local / 2;
158171
std::vector<float*> to_global;
159172
for (size_t i = 0; i < drain_count && !local.empty(); ++i) {
160173
to_global.push_back(local.back());
@@ -163,6 +176,17 @@ struct ThreadCache {
163176
global->release_batch(to_global, cls);
164177
local.push_back(ptr);
165178
}
179+
180+
// Free all cached buffers in this thread's cache directly (bypass global).
181+
void flush() {
182+
for (auto& kv : buckets) {
183+
for (float* ptr : kv.second) {
184+
if (g_free_fn) g_free_fn(ptr);
185+
else std::free(ptr);
186+
}
187+
kv.second.clear();
188+
}
189+
}
166190
};
167191

168192
thread_local ThreadCache t_thread_cache;
@@ -187,15 +211,34 @@ void MemoryPool::release(float* ptr, size_t original_n) {
187211
}
188212

189213
void MemoryPool::trim() {
214+
// Flush calling thread's local cache first (frees directly, no lock needed).
215+
memory_pool_detail::t_thread_cache.flush();
216+
217+
// Then free global pool buckets.
190218
auto* p = impl();
191219
std::lock_guard<std::mutex> lock(p->mutex);
192-
for (auto& [cls, free_list] : p->buckets) {
193-
for (float* ptr : free_list) {
220+
for (auto& kv : p->buckets) {
221+
for (float* ptr : kv.second) {
194222
if (g_free_fn) g_free_fn(ptr);
195223
else std::free(ptr);
196224
}
197-
free_list.clear();
225+
kv.second.clear();
226+
}
227+
}
228+
229+
size_t MemoryPool::cached_bytes() const {
230+
size_t total = 0;
231+
// Thread-local cache
232+
for (auto& kv : memory_pool_detail::t_thread_cache.buckets) {
233+
total += kv.first * sizeof(float) * kv.second.size();
234+
}
235+
// Global pool
236+
auto* p = const_cast<MemoryPool*>(this)->impl();
237+
std::lock_guard<std::mutex> lock(p->mutex);
238+
for (auto& kv : p->buckets) {
239+
total += kv.first * sizeof(float) * kv.second.size();
198240
}
241+
return total;
199242
}
200243

201244
std::shared_ptr<float> MemoryPool::acquire_shared(size_t n) {

‎core/memory_pool.h‎

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -28,10 +28,14 @@ class MemoryPool {
2828
// Round requested size up to the pool's size class (for tests/tooling).
2929
static size_t size_class(size_t n);
3030

31-
// Free all cached buffers in the global pool (does not affect in-use buffers).
31+
// Free all cached buffers in the global pool AND calling thread's local cache.
3232
// Call periodically to prevent unbounded memory growth.
3333
void trim();
3434

35+
// Return total bytes currently cached (global + calling thread's local).
36+
// For diagnostics only — approximate under concurrency.
37+
size_t cached_bytes() const;
38+
3539
// Custom allocator hooks (for CUDA managed memory).
3640
// When set, these replace std::malloc/std::free for all pool allocations.
3741
using AllocFn = float*(*)(size_t n_bytes);

‎core/tensor.cpp‎

Lines changed: 24 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -20,12 +20,26 @@
2020

2121
static thread_local std::mt19937 rng(42);
2222

23-
Tensor::Tensor() : requires_grad(false), data_size_(0), grad_size_(0) {}
23+
// Global live tensor counter (atomic for thread safety).
24+
static std::atomic<int64_t> g_live_tensor_count{0};
25+
static std::atomic<int64_t> g_live_tensor_bytes{0};
2426

25-
Tensor::~Tensor() {}
27+
int64_t Tensor::live_count() { return g_live_tensor_count.load(std::memory_order_relaxed); }
28+
int64_t Tensor::live_bytes() { return g_live_tensor_bytes.load(std::memory_order_relaxed); }
29+
30+
Tensor::Tensor() : requires_grad(false), data_size_(0), grad_size_(0) {
31+
g_live_tensor_count.fetch_add(1, std::memory_order_relaxed);
32+
}
33+
34+
Tensor::~Tensor() {
35+
g_live_tensor_count.fetch_sub(1, std::memory_order_relaxed);
36+
g_live_tensor_bytes.fetch_sub(static_cast<int64_t>(data_size_ + grad_size_) * sizeof(float),
37+
std::memory_order_relaxed);
38+
}
2639

2740
Tensor::Tensor(const std::vector<size_t>& shape, bool requires_grad)
2841
: shape(shape), requires_grad(requires_grad) {
42+
g_live_tensor_count.fetch_add(1, std::memory_order_relaxed);
2943
size_t total = 1;
3044
for (auto s : shape) total *= s;
3145
data_size_ = total;
@@ -40,10 +54,13 @@ Tensor::Tensor(const std::vector<size_t>& shape, bool requires_grad)
4054
} else {
4155
grad_size_ = 0;
4256
}
57+
g_live_tensor_bytes.fetch_add(static_cast<int64_t>(data_size_ + grad_size_) * sizeof(float),
58+
std::memory_order_relaxed);
4359
}
4460

4561
Tensor::Tensor(const std::vector<float>& data, const std::vector<size_t>& shape, bool requires_grad)
4662
: shape(shape), requires_grad(requires_grad) {
63+
g_live_tensor_count.fetch_add(1, std::memory_order_relaxed);
4764
data_size_ = data.size();
4865
data_storage_ = MemoryPool::instance().acquire_shared(data_size_);
4966
if (!data_storage_) throw std::bad_alloc();
@@ -56,12 +73,15 @@ Tensor::Tensor(const std::vector<float>& data, const std::vector<size_t>& shape,
5673
} else {
5774
grad_size_ = 0;
5875
}
76+
g_live_tensor_bytes.fetch_add(static_cast<int64_t>(data_size_ + grad_size_) * sizeof(float),
77+
std::memory_order_relaxed);
5978
}
6079

6180
Tensor::Tensor(std::shared_ptr<float> data_storage, size_t data_size,
6281
const std::vector<size_t>& shape, bool requires_grad)
6382
: shape(shape), requires_grad(requires_grad),
6483
data_storage_(std::move(data_storage)), data_size_(data_size) {
84+
g_live_tensor_count.fetch_add(1, std::memory_order_relaxed);
6585
if (requires_grad) {
6686
grad_size_ = data_size_;
6787
grad_storage_ = MemoryPool::instance().acquire_shared(grad_size_);
@@ -70,6 +90,8 @@ Tensor::Tensor(std::shared_ptr<float> data_storage, size_t data_size,
7090
} else {
7191
grad_size_ = 0;
7292
}
93+
g_live_tensor_bytes.fetch_add(static_cast<int64_t>(data_size_ + grad_size_) * sizeof(float),
94+
std::memory_order_relaxed);
7395
}
7496

7597
TensorPtr Tensor::create(const std::vector<size_t>& shape, bool requires_grad) {

‎core/tensor.h‎

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@
1111
#include <cstdint>
1212
#include <cstdio>
1313
#include <random>
14+
#include <atomic>
1415

1516
enum class DType { Float32, Float16 };
1617

@@ -172,6 +173,10 @@ class Tensor : public std::enable_shared_from_this<Tensor> {
172173

173174
void print(const char* name = nullptr) const;
174175

176+
// Diagnostics: live tensor count/memory (atomic, thread-safe).
177+
static int64_t live_count();
178+
static int64_t live_bytes();
179+
175180
private:
176181
void build_topo(std::vector<Tensor*>& topo, std::vector<Tensor*>& visited);
177182
bool should_track_grad() const;

‎examples/resnet18_imagenette.cpp‎

Lines changed: 53 additions & 42 deletions
Original file line numberDiff line numberDiff line change
@@ -647,63 +647,74 @@ int main(int argc, char* argv[]) {
647647
size_t batch_end = std::min(batch_start + batch_size, n_train);
648648
size_t bs = batch_end - batch_start;
649649

650-
// Assemble batch from mmap using shuffled indices
651-
auto images = Tensor::create({bs, 3, 224, 224}, false);
652-
auto labels = Tensor::create({bs}, false);
653-
654-
for (size_t i = 0; i < bs; i++) {
655-
size_t idx = train_indices[batch_start + i];
656-
std::memcpy(images->data() + i * img_elems,
657-
train_images_mmap.data + idx * img_elems,
658-
img_elems * sizeof(float));
659-
labels->data()[i] = train_labels_mmap.data[idx];
660-
}
650+
float batch_loss = 0.0f;
651+
652+
// Scoped block: all batch tensors freed before trim()
653+
{
654+
// Assemble batch from mmap using shuffled indices
655+
auto images = Tensor::create({bs, 3, 224, 224}, false);
656+
auto labels = Tensor::create({bs}, false);
657+
658+
for (size_t i = 0; i < bs; i++) {
659+
size_t idx = train_indices[batch_start + i];
660+
std::memcpy(images->data() + i * img_elems,
661+
train_images_mmap.data + idx * img_elems,
662+
img_elems * sizeof(float));
663+
labels->data()[i] = train_labels_mmap.data[idx];
664+
}
661665

662-
// Data augmentation: pad 16 -> random crop 224x224 -> random horizontal flip
663-
auto augmented = images->pad2d(16)->random_crop(224, 224)->random_flip_horizontal(0.5f);
666+
// Data augmentation: pad 16 -> random crop 224x224 -> random horizontal flip
667+
auto augmented = images->pad2d(16)->random_crop(224, 224)->random_flip_horizontal(0.5f);
664668

665-
optimizer.zero_grad();
669+
optimizer.zero_grad();
666670

667-
auto output = model.forward(augmented);
668-
auto loss = criterion(output, labels);
671+
auto output = model.forward(augmented);
672+
auto loss = criterion(output, labels);
669673

670-
loss->backward();
674+
loss->backward();
671675

672-
// Manual L2 weight decay on conv/linear weights
673-
apply_weight_decay(all_params, weight_decay);
676+
// Manual L2 weight decay on conv/linear weights
677+
apply_weight_decay(all_params, weight_decay);
674678

675-
optimizer.step();
679+
optimizer.step();
676680

677-
// Free cached memory to prevent unbounded growth with large 224x224 tensors
678-
MemoryPool::instance().trim();
679-
680-
// Read loss scalar (handles GPU -> CPU transfer)
681-
total_loss += read_loss_scalar(loss);
682-
num_batches++;
681+
// Extract loss and accuracy BEFORE releasing tensors
682+
batch_loss = read_loss_scalar(loss);
683683

684-
// Compute training accuracy for this batch
685-
auto cpu_output = ensure_cpu(output);
686-
auto cpu_labels = ensure_cpu(labels);
687-
size_t nc = cpu_output->shape[1];
688-
for (size_t i = 0; i < bs; i++) {
689-
size_t pred = 0;
690-
float mv = cpu_output->data()[i * nc];
691-
for (size_t j = 1; j < nc; j++) {
692-
if (cpu_output->data()[i * nc + j] > mv) {
693-
mv = cpu_output->data()[i * nc + j];
694-
pred = j;
684+
auto cpu_output = ensure_cpu(output);
685+
auto cpu_labels = ensure_cpu(labels);
686+
size_t nc = cpu_output->shape[1];
687+
for (size_t i = 0; i < bs; i++) {
688+
size_t pred = 0;
689+
float mv = cpu_output->data()[i * nc];
690+
for (size_t j = 1; j < nc; j++) {
691+
if (cpu_output->data()[i * nc + j] > mv) {
692+
mv = cpu_output->data()[i * nc + j];
693+
pred = j;
694+
}
695695
}
696+
if (pred == static_cast<size_t>(cpu_labels->data()[i])) correct++;
697+
total++;
696698
}
697-
if (pred == static_cast<size_t>(cpu_labels->data()[i])) correct++;
698-
total++;
699699
}
700+
// All batch tensors (images, labels, augmented, output, loss) now freed.
701+
702+
// Free cached memory — thread-local + global pools.
703+
MemoryPool::instance().trim();
704+
705+
total_loss += batch_loss;
706+
num_batches++;
700707

701-
// Print progress
708+
// Print progress + memory diagnostics
702709
if (epoch == start_epoch || num_batches % 20 == 0) {
703710
auto batch_time = std::chrono::high_resolution_clock::now();
704711
double elapsed = std::chrono::duration<double>(batch_time - epoch_start).count();
705-
printf("\r Epoch %3d: batch %3zu/%zu | %.1fs elapsed",
706-
epoch + 1, num_batches, total_batches, elapsed);
712+
int64_t live = Tensor::live_count();
713+
int64_t live_mb = Tensor::live_bytes() / (1024 * 1024);
714+
size_t pool_mb = MemoryPool::instance().cached_bytes() / (1024 * 1024);
715+
printf("\r Epoch %3d: batch %3zu/%zu | %.1fs | tensors: %lld (%lldMB) pool: %zuMB",
716+
epoch + 1, num_batches, total_batches, elapsed,
717+
(long long)live, (long long)live_mb, pool_mb);
707718
fflush(stdout);
708719
}
709720
}

0 commit comments

Comments
 (0)