@@ -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