diff -u b/term_reference_tree.module b/term_reference_tree.module
--- b/term_reference_tree.module
+++ b/term_reference_tree.module
@@ -86,35 +86,33 @@
  * @return array
  *   A nested array of the term's child objects.
  */
-function _term_reference_tree_get_term_hierarchy($tid, $vid, $allowed = array(), $default = array(), $opened = array(), $use_ajax = FALSE, $max_depth = NULL) {
+function _term_reference_tree_get_term_hierarchy($tid, $vid, $allowed = array(), $expanded = array(), $use_ajax = FALSE, $max_depth = NULL) {
   $tree = _term_reference_tree_taxonomy_get_tree($vid);
   if (!empty($allowed)) {
     $tree['terms'] = array_intersect_key($tree['terms'], $allowed);
   }
 
-  $default_parents = _term_reference_tree_taxonomy_get_parents_all($vid, $default);
-  return _term_reference_tree_get_term_hierarchy_recursive($tid, $tree, $default_parents, $opened, $use_ajax, $max_depth);
+  return _term_reference_tree_get_term_hierarchy_recursive($tid, $tree, $expanded, $use_ajax, $max_depth);
 }
 
 /**
  * Recursive helper function for _term_reference_tree_get_term_hierarchy().
  */
-function _term_reference_tree_get_term_hierarchy_recursive($parent_tid, $tree, $default, $opened, $use_ajax, $max_depth, $depth = 1) {
+function _term_reference_tree_get_term_hierarchy_recursive($parent_tid, $tree, $expanded, $use_ajax, $max_depth, $depth = 1) {
   $terms = array();
 
   if (isset($tree['children'][$parent_tid])) {
     foreach($tree['children'][$parent_tid] as $child_tid) {
       $term = $tree['terms'][$child_tid];
-      $max_depth_reached = isset($max_depth) && $depth >= $max_depth;
-      $term->has_children = isset($tree['children'][$term->tid]) && !$max_depth_reached;
+      $term->has_children = isset($tree['children'][$term->tid]);
 
       if ($term->has_children) {
         // Process children if:
-        // - we don't use ajax,
-        // - OR if this term is opened via ajax,
-        // - OR if one of the children of this term is in $default.
-        if (!$use_ajax || in_array($term->tid, $opened) || in_array($term->tid, $default)) {
-          $term->children = _term_reference_tree_get_term_hierarchy_recursive($term->tid, $tree, $default, $opened, $use_ajax, $max_depth, ++$depth);
+        // Max depth is not reached
+        // And we don't use ajax OR if this term must be expanded,
+        $max_depth_reached = isset($max_depth) && $depth >= $max_depth;
+        if (!$max_depth_reached && (!$use_ajax || in_array($term->tid, $expanded))) {
+          $term->children = _term_reference_tree_get_term_hierarchy_recursive($term->tid, $tree, $expanded, $use_ajax, $max_depth, $depth + 1);
         }
       }
 
diff -u b/term_reference_tree.widget.inc b/term_reference_tree.widget.inc
--- b/term_reference_tree.widget.inc
+++ b/term_reference_tree.widget.inc
@@ -289,7 +289,9 @@
     $max_depth = !empty($element['#max_depth']) ? $element['#max_depth'] : NULL;
 
     // Build the terms hierarchy.
-    $element['#options_tree'] = _term_reference_tree_get_term_hierarchy($element['#parent_tid'], $element['#vocabulary']->vid, $allowed, $default, $opened, $element['#use_ajax'], $max_depth);
+    $default_parents = _term_reference_tree_taxonomy_get_parents_all($element['#vocabulary']->vid, $default);
+    $element['#expanded'] = $opened + $default_parents;
+    $element['#options_tree'] = _term_reference_tree_get_term_hierarchy($element['#parent_tid'], $element['#vocabulary']->vid, $allowed, $element['#expanded'], $element['#use_ajax'], $max_depth);
 
     // Flatten the hierarchy.
     $terms_flat = _term_reference_tree_hierarchy_flatten($element['#options_tree']);
@@ -727,7 +729,8 @@
 
   $container[$term->tid] = $e;
 
-  if ($term->has_children) {
+  $max_depth_reached = !empty($element['#max_depth']) && $depth >= $element['#max_depth'];
+  if ($term->has_children && !$max_depth_reached) {
     $container['#has_children'] = TRUE;
     $parents = $parent_tids;
     $parents[] = $term->tid;
@@ -760,11 +763,12 @@
  *   A completed checkbox_tree_level element.
  */
 function _term_reference_tree_build_level($element, $term, $form_state, $value, $max_choices, $parent_tids, $depth) {
+  $must_expand = (!$element['#start_minimized'] && !$element['#use_ajax']) || (isset($term->tid) && in_array($term->tid, $element['#expanded']));
   $container = array(
     '#max_choices' => $max_choices,
     '#leaves_only' => isset($element['#leaves_only']) ? $element['#leaves_only'] : FALSE,
     '#start_minimized' => isset($element['#start_minimized']) ? $element['#start_minimized'] : FALSE,
-    '#level_start_minimized' => $depth > 1 &&  $element['#start_minimized'] && empty($term->children),
+    '#level_start_minimized' => $depth > 1 && !$must_expand,
     '#depth' => $depth,
   );
 
