Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions src/cmlibs/utils/zinc/field.py
Original file line number Diff line number Diff line change
Expand Up @@ -554,7 +554,7 @@ def create_field_stored_mesh_location(fieldmodule: Fieldmodule, mesh: Mesh, name

:param fieldmodule: Zinc fieldmodule to find or create field in.
:param mesh: Mesh to store locations in, from same fieldmodule.
:param name: Name of new field. If not defined, defaults to "location\_" + mesh.getName().
:param name: Name of new field. If not defined, defaults to "location_" + mesh.getName().
:param managed: Managed state of field.
:return: Zinc FieldStoredMeshLocation
"""
Expand All @@ -576,7 +576,7 @@ def find_or_create_field_stored_mesh_location(fieldmodule: Fieldmodule, mesh: Me

:param fieldmodule: Zinc fieldmodule to find or create field in.
:param mesh: Mesh to store locations in, from same fieldmodule.
:param name: Name of new field. If not defined, defaults to "location\_" + mesh.getName().
:param name: Name of new field. If not defined, defaults to "location_" + mesh.getName().
:param managed: Managed state of field if created here.
"""
if not name:
Expand Down
3 changes: 2 additions & 1 deletion src/cmlibs/utils/zinc/general.py
Original file line number Diff line number Diff line change
Expand Up @@ -124,7 +124,7 @@ def create_node(field_module, data_object, identifier=-1, node_set_name='nodes',
"""
Create a Node in the field_module using the data_object. The data object must supply a 'get_field_names' method
and a 'get_time_sequence' method. Derive a node data object from the 'AbstractNodeDataObject' class to ensure
that the data object class meets it's requirements.
that the data object class meets its requirements.

Optionally use the identifier to set the identifier of the Node created, the time parameter to set
the time value in the cache, or the node_set_name to specify which node set to use the default node set
Expand Down Expand Up @@ -162,6 +162,7 @@ def create_node(field_module, data_object, identifier=-1, node_set_name='nodes',
for i, field in enumerate(fields):
field_name = field_names[i]
field_value = getattr(data_object, field_name)()
# print(field_name, type(field_value), field_value, isinstance(field_value, ("".__class__, u"".__class__)))
if isinstance(field_value, ("".__class__, u"".__class__)):
field.assignString(field_cache, field_value)
else:
Expand Down
8 changes: 8 additions & 0 deletions src/cmlibs/utils/zinc/group.py
Original file line number Diff line number Diff line change
Expand Up @@ -503,20 +503,25 @@ def match_fitting_group_names(data_fieldmodule, model_fieldmodule, log_diagnosti
:param data_fieldmodule: Data Fieldmodule whose group names may be modified.
:param model_fieldmodule: Model Fieldmodule containing preferred group names.
:param log_diagnostics: Set to True to write diagonstic messages about name matches and changes to logging.
:return: A dictionary of matched group names.
"""
# future: match with annotation terms
model_names = [group.getName() for group in get_group_list(model_fieldmodule)]
matched_names = {}
for data_group in get_group_list(data_fieldmodule):
data_name = data_group.getName()
compare_name = data_name.strip().casefold()
for model_name in model_names:
# print('comparing "%s" to "%s" or "%s"' % (data_name, model_name, compare_name))
if model_name == data_name:
matched_names[model_name] = (data_name, None)
if log_diagnostics:
logger.info("Data group '" + data_name + "' found in model")
break
elif model_name.strip().casefold() == compare_name:
result = data_group.setName(model_name)
if result == RESULT_OK:
matched_names[model_name] = (data_name, compare_name)
if log_diagnostics:
logger.info("Data group '" + data_name + "' found in model as '" +
model_name + "'. Renaming to match.")
Expand All @@ -527,5 +532,8 @@ def match_fitting_group_names(data_fieldmodule, model_fieldmodule, log_diagnosti
logger.error(" Reason: field of that name already exists.")
break
else:
# print('Data group "' + data_name + '" not found in model')
if log_diagnostics:
logger.info("Data group '" + data_name + "' not found in model")

return matched_names
21 changes: 13 additions & 8 deletions src/cmlibs/utils/zinc/region.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,18 +11,19 @@ def _find_missing(lst):


def convert_nodes_to_datapoints(target_region, source_region, source_nodeset_type=Field.DOMAIN_TYPE_NODES,
destroy_after_conversion=True):
destroy_after_conversion=True, field_names=None):
"""
When the source nodeset type is Field.DOMAIN_TYPE_DATAPOINTS, then datapoints are transferred from the
Converts nodes in the source region to datapoints in the target region, renumbering any existing
datapoints in target region to not clash.
When the source nodeset type is Field.DOMAIN_TYPE_DATAPOINTS, then datapoints are transferred from the
source region to the target region.
:param target_region: Zinc Region to read data into. Existing data points are renumbered to avoid nodes.
:param source_region: Zinc Region containing nodes to transfer.
:param source_nodeset_type: Set to Field.DOMAIN_TYPE_DATAPOINTS or Field.DOMAIN_TYPE_NODES to transfer datapoints
:param source_nodeset_type: Set to Field.DOMAIN_TYPE_DATAPOINTS or Field.DOMAIN_TYPE_NODES to transfer datapoints
or convert nodes. Datapoint transfer should only be to different regions [default: Field.DOMAIN_TYPE_NODES].
:param destroy_after_conversion: Set to True to destroy nodes that have been successfully converted, or False
:param destroy_after_conversion: Set to True to destroy nodes that have been successfully converted, or False
to leave intact in source region [default: True].
:param field_names: A list of field names to output.
"""
source_fieldmodule = source_region.getFieldmodule()
target_fieldmodule = target_region.getFieldmodule()
Expand All @@ -46,19 +47,22 @@ def convert_nodes_to_datapoints(target_region, source_region, source_nodeset_typ
datapoint_identifier = datapoint.getIdentifier()
if datapoint_identifier in existing_nodes_identifiers_set and len(available_identifiers):
next_identifier = available_identifiers.pop(0)
elif datapoint_identifier not in existing_nodes_identifiers_set:
next_identifier = datapoint_identifier
else:
max_identifier += 1
next_identifier = max_identifier

identifier_map[datapoint_identifier] = next_identifier
if next_identifier != datapoint_identifier:
identifier_map[datapoint_identifier] = next_identifier
datapoint = datapoint_iterator.next()

for current_identifier, new_identifier in identifier_map.items():
datapoint = datapoints.findNodeByIdentifier(current_identifier)
datapoint.setIdentifier(new_identifier)

# transfer nodes as datapoints to target_region
buffer = write_to_buffer(source_region, resource_domain_type=source_nodeset_type)
buffer = write_to_buffer(source_region, resource_domain_type=source_nodeset_type, field_names=field_names)
if source_nodeset_type == Field.DOMAIN_TYPE_NODES:
buffer = buffer.replace(bytes("!#nodeset nodes", "utf-8"), bytes("!#nodeset datapoints", "utf-8"))
result = read_from_buffer(target_region, buffer)
Expand All @@ -68,17 +72,18 @@ def convert_nodes_to_datapoints(target_region, source_region, source_nodeset_typ
nodes.destroyAllNodes()


def copy_fitting_data(target_region, source_region):
def copy_fitting_data(target_region, source_region, field_names=None):
"""
Copy nodes and data points from source_region to target_region, converting nodes to data points and
offsetting data point identifiers to not clash. All groups and fields in use are transferred.
This is used for setting up fitting problems where data needs to be in datapoints only.
:param target_region: Zinc Region to read nodes/data into.
:param source_region: Zinc Region containing nodes/data to transfer. Unmodified.
:param field_names: A list of field names to output.
"""
for domain_type in [Field.DOMAIN_TYPE_DATAPOINTS, Field.DOMAIN_TYPE_NODES]:
convert_nodes_to_datapoints(target_region, source_region, source_nodeset_type=domain_type,
destroy_after_conversion=False)
destroy_after_conversion=False, field_names=field_names)


def copy_nodeset(region, nodeset):
Expand Down
35 changes: 31 additions & 4 deletions tests/test_zinc_group.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,3 @@
import os
import unittest
from cmlibs.utils.zinc.field import find_or_create_field_group
from cmlibs.utils.zinc.group import (
Expand All @@ -8,6 +7,7 @@
from cmlibs.zinc.element import Element
from cmlibs.zinc.field import Field
from cmlibs.zinc.result import RESULT_OK

from utilities import assert_almost_equal_list, get_test_resource_name


Expand Down Expand Up @@ -136,7 +136,6 @@ def test_group_add_compare_group_local_contents(self):
self.assertEqual(0, group1.getMeshGroup(mesh1d).getSize())
self.assertEqual(0, group1.getNodesetGroup(nodes).getSize())


def test_match_fitting_group_names(self):
"""
Test utility functions for adding and comparing group local contents.
Expand All @@ -154,8 +153,36 @@ def test_match_fitting_group_names(self):
data_group_fred = find_or_create_field_group(data_fieldmodule, " fRed\t")
data_group_two_names = find_or_create_field_group(data_fieldmodule, "\t two NAMES ")

match_fitting_group_names(data_fieldmodule, model_fieldmodule, log_diagnostics=True)

names = match_fitting_group_names(data_fieldmodule, model_fieldmodule, log_diagnostics=True)
self.assertEqual(data_group_bob.getName(), "bob")
self.assertEqual(data_group_fred.getName(), "fred")
self.assertEqual(data_group_two_names.getName(), "two names")
self.assertIn("two names", names)
self.assertEqual(names["two names"], ('\t two NAMES ', 'two names'))
self.assertIn("bob", names)
self.assertEqual(names["bob"], (' Bob', 'bob'))

def test_match_fitting_group_names_cap(self):
"""
Test utility functions for adding and comparing group local contents.
"""
context = Context("test")
model_region = context.createRegion()
model_fieldmodule = model_region.getFieldmodule()
find_or_create_field_group(model_fieldmodule, "Bob", managed=True)
find_or_create_field_group(model_fieldmodule, "fRed", managed=True)
find_or_create_field_group(model_fieldmodule, "james", managed=True)

data_region = context.createRegion()
data_fieldmodule = data_region.getFieldmodule()
data_group_bob = find_or_create_field_group(data_fieldmodule, "bob")
data_group_fred = find_or_create_field_group(data_fieldmodule, "fred")
find_or_create_field_group(data_fieldmodule, "james")

names = match_fitting_group_names(data_fieldmodule, model_fieldmodule, log_diagnostics=True)
self.assertEqual(data_group_bob.getName(), "Bob")
self.assertEqual(data_group_fred.getName(), "fRed")
self.assertIn("Bob", names)
self.assertEqual(names["Bob"], ('bob', 'bob'))
self.assertIn("james", names)
self.assertEqual(names["james"], ('james', None))
Loading
Loading