Skip to content
Merged
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
6 changes: 6 additions & 0 deletions pyunitwizard/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,3 +41,9 @@ def __print_version__():
except:
pass

try:
import astropy.units # noqa: F401
configure.load_library('astropy.units')
except:
pass

2 changes: 1 addition & 1 deletion pyunitwizard/_private/forms.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
from pyunitwizard import kernel

forms = ['openmm.unit', 'pint', 'unyt', 'string']
forms = ['openmm.unit', 'pint', 'unyt', 'astropy.units', 'string']

def digest_form(form: str) -> str:
""" Check if the form is correct.
Expand Down
2 changes: 1 addition & 1 deletion pyunitwizard/_private/parsers.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
parsers = ['openmm.unit', 'pint', 'unyt']
parsers = ['openmm.unit', 'pint', 'unyt', 'astropy.units']

def digest_parser(parser: str) -> str:
""" Check if parser is correct."""
Expand Down
12 changes: 7 additions & 5 deletions pyunitwizard/configure/configure.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,12 +8,14 @@
from importlib.util import find_spec
from typing import List, Dict, Union

libraries = ['pint', 'openmm.unit', 'unyt']
parsers = ['pint', 'openmm.unit', 'unyt']
libraries = ['pint', 'openmm.unit', 'unyt', 'astropy.units']
parsers = ['pint', 'openmm.unit', 'unyt', 'astropy.units']
_aux_dict_modules = {
'pint':'pint',
'openmm.unit':'openmm',
'unyt': 'unyt'}
'pint': 'pint',
'openmm.unit': 'openmm',
'unyt': 'unyt',
'astropy.units': 'astropy',
}

def reset() -> None:
"""Resets all kernel variables."""
Expand Down
7 changes: 6 additions & 1 deletion pyunitwizard/forms/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,12 @@
dict_compatibility={}

_base_package = __name__.replace('.base','')
_forms_apis_modules = {'openmm.unit':'api_openmm_unit', 'pint':'api_pint', 'unyt':'api_unyt'}
_forms_apis_modules = {
'openmm.unit': 'api_openmm_unit',
'pint': 'api_pint',
'unyt': 'api_unyt',
'astropy.units': 'api_astropy_unit',
}

def load_library(library: str) -> None:
""" Loads a library. This means that it updates all dictionaries defined above
Expand Down
162 changes: 162 additions & 0 deletions pyunitwizard/forms/api_astropy_unit.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,162 @@
from typing import Any, Dict, Union

from pyunitwizard._private.exceptions import LibraryNotFoundError
from pyunitwizard._private.quantity_or_unit import ArrayLike

try:
from astropy import units as astropy_units
except Exception as exc: # pragma: no cover - handled through exception
raise LibraryNotFoundError('astropy') from exc

AstropyQuantity = astropy_units.Quantity
AstropyUnitBase = astropy_units.UnitBase

form_name = 'astropy.units'
parser = True

is_form = {
AstropyQuantity: form_name,
AstropyUnitBase: form_name,
}


def _to_unit(quantity_or_unit: Union[AstropyQuantity, AstropyUnitBase]) -> AstropyUnitBase:
if is_quantity(quantity_or_unit):
return get_unit(quantity_or_unit)
if is_unit(quantity_or_unit):
return quantity_or_unit
raise TypeError('Expected an astropy quantity or unit')


def is_quantity(quantity_or_unit: Any) -> bool:
return isinstance(quantity_or_unit, AstropyQuantity)


def is_unit(quantity_or_unit: Any) -> bool:
return isinstance(quantity_or_unit, AstropyUnitBase)


_dimensions_translator = {
'm': '[L]',
'kg': '[M]',
's': '[T]',
'K': '[K]',
'mol': '[mol]',
'A': '[A]',
'cd': '[Cd]',
}


def dimensionality(quantity_or_unit: Union[AstropyQuantity, AstropyUnitBase]) -> Dict[str, float]:
unit = _to_unit(quantity_or_unit)
decomposed = unit.decompose(bases=astropy_units.si.bases)
dimensionality_dict = {'[L]': 0, '[M]': 0, '[T]': 0, '[K]': 0, '[mol]': 0, '[A]': 0, '[Cd]': 0}
Comment thread
dprada marked this conversation as resolved.

for base, power in zip(decomposed.bases, decomposed.powers):
key = _dimensions_translator.get(base.to_string())
if key is None:
raise ValueError(f"Unrecognized base unit: {base.to_string()} in {unit}")
dimensionality_dict[key] += float(power)

return dimensionality_dict


def compatibility(quantity_or_unit_1: Union[AstropyQuantity, AstropyUnitBase],
quantity_or_unit_2: Union[AstropyQuantity, AstropyUnitBase]) -> bool:
unit_1 = _to_unit(quantity_or_unit_1)
unit_2 = _to_unit(quantity_or_unit_2)
return unit_1.is_equivalent(unit_2)


def make_quantity(value: Union[int, float, ArrayLike],
unit: Union[str, AstropyUnitBase]) -> AstropyQuantity:
unit_obj = astropy_units.Unit(unit)
return astropy_units.Quantity(value, unit_obj)


def get_value(quantity: AstropyQuantity) -> Union[int, float, ArrayLike]:
return quantity.value


def get_unit(quantity: AstropyQuantity) -> AstropyUnitBase:
return quantity.unit


def change_value(quantity: AstropyQuantity,
value: Union[int, float, ArrayLike]) -> AstropyQuantity:
return make_quantity(value, get_unit(quantity))


def convert(quantity: AstropyQuantity,
unit: Union[str, AstropyUnitBase]) -> AstropyQuantity:
unit_obj = astropy_units.Unit(unit)
return quantity.to(unit_obj)


# Parser

def string_to_quantity(string: str) -> AstropyQuantity:
return astropy_units.Quantity(string)


def string_to_unit(string: str) -> AstropyUnitBase:
return astropy_units.Unit(string)


# To string

def quantity_to_string(quantity: AstropyQuantity) -> str:
return str(quantity)


def unit_to_string(unit: AstropyUnitBase) -> str:
return unit.to_string()


# To pint

def quantity_to_pint(quantity: AstropyQuantity):
from .api_pint import make_quantity as make_pint_quantity

value = get_value(quantity)
unit_name = unit_to_string(get_unit(quantity))
return make_pint_quantity(value, unit_name)


def unit_to_pint(unit: AstropyUnitBase):
from .api_pint import get_unit as get_pint_unit

quantity = quantity_to_pint(1.0 * unit)
return get_pint_unit(quantity)


# To openmm.unit

def quantity_to_openmm_unit(quantity: AstropyQuantity):
from .api_pint import quantity_to_openmm_unit as pint_to_openmm_unit

pint_quantity = quantity_to_pint(quantity)
return pint_to_openmm_unit(pint_quantity)


def unit_to_openmm_unit(unit: AstropyUnitBase):
from .api_openmm_unit import get_unit as get_openmm_unit

quantity = quantity_to_openmm_unit(1.0 * unit)
return get_openmm_unit(quantity)


# To unyt

def quantity_to_unyt(quantity: AstropyQuantity):
from .api_pint import quantity_to_unyt as pint_to_unyt

pint_quantity = quantity_to_pint(quantity)
return pint_to_unyt(pint_quantity)


def unit_to_unyt(unit: AstropyUnitBase):
from .api_unyt import get_unit as get_unyt_unit

quantity = quantity_to_unyt(1.0 * unit)
return get_unyt_unit(quantity)
26 changes: 24 additions & 2 deletions pyunitwizard/forms/api_openmm_unit.py
Original file line number Diff line number Diff line change
Expand Up @@ -306,12 +306,12 @@ def quantity_to_unyt(quantity: openmm_unit.Quantity):

def unit_to_unyt(unit: openmm_unit.Unit):
""" Transform a unit from openmm.unit to a unyt unit.

Parameters
-----------
unit : openmm.unit.Unit
A unit.

Returns
-------
unyt_unit
Expand All @@ -323,3 +323,25 @@ def unit_to_unyt(unit: openmm_unit.Unit):

return get_unyt_unit(quantity)


## To astropy.units

def quantity_to_astropy_units(quantity: openmm_unit.Quantity):
""" Transform a quantity from openmm.unit to astropy.units."""

from .api_pint import quantity_to_astropy_units as pint_to_astropy_units

pint_quantity = quantity_to_pint(quantity)

return pint_to_astropy_units(pint_quantity)


def unit_to_astropy_units(unit: openmm_unit.Unit):
""" Transform a unit from openmm.unit to astropy.units."""

from .api_astropy_unit import get_unit as get_astropy_unit

quantity = quantity_to_astropy_units(1.0*unit)

return get_astropy_unit(quantity)

27 changes: 25 additions & 2 deletions pyunitwizard/forms/api_pint.py
Original file line number Diff line number Diff line change
Expand Up @@ -331,12 +331,12 @@ def quantity_to_unyt(quantity: pint.Quantity):

def unit_to_unyt(unit: pint.Unit):
""" Transform a unit from a pint unit to a unyt unit.

Parameters
-----------
unit : pint.Unit
A unit.

Returns
-------
unyt_array or unyt_quantity
Expand All @@ -350,3 +350,26 @@ def unit_to_unyt(unit: pint.Unit):
return get_unyt_unit(quantity)


## To astropy.units

def quantity_to_astropy_units(quantity: pint.Quantity):
""" Transform a quantity from pint to astropy.units."""

from .api_astropy_unit import make_quantity as make_astropy_quantity

value = get_value(quantity)
unit_name = unit_to_string(get_unit(quantity))

return make_astropy_quantity(value, unit_name)


def unit_to_astropy_units(unit: pint.Unit):
""" Transform a unit from pint to astropy.units."""

from .api_astropy_unit import get_unit as get_astropy_unit

quantity = quantity_to_astropy_units(1.0*unit)

return get_astropy_unit(quantity)


20 changes: 18 additions & 2 deletions pyunitwizard/forms/api_string.py
Original file line number Diff line number Diff line change
Expand Up @@ -234,12 +234,12 @@ def quantity_to_pint(quantity: str):

def unit_to_pint(unit: str):
""" Transform a quantity from a string quantity to a pint quantity.

Parameters
-----------
quantity : str
A quanitity.

Returns
-------
pint.Quantity
Expand All @@ -260,3 +260,19 @@ def quantity_to_unyt(quantity: str):
def unit_to_unyt(quantity: str):
raise NotImplementedError


## To astropy.units

def quantity_to_astropy_units(quantity: str):
from .api_astropy_unit import string_to_quantity as _string_to_quantity

return _string_to_quantity(quantity)


def unit_to_astropy_units(unit: str):
from .api_astropy_unit import get_unit as get_astropy_unit

quantity = quantity_to_astropy_units(unit)

return get_astropy_unit(quantity)

26 changes: 24 additions & 2 deletions pyunitwizard/forms/api_unyt.py
Original file line number Diff line number Diff line change
Expand Up @@ -295,12 +295,12 @@ def quantity_to_openmm_unit(quantity: Union[unyt_array, unyt_quantity]):

def unit_to_openmm_unit(unit: unyt_unit):
""" Transform a unit from unyt to a openmm.unit unit.

Parameters
-----------
unit : unyt_unit
A unit.

Returns
-------
openmm_unit.Unit
Expand All @@ -313,3 +313,25 @@ def unit_to_openmm_unit(unit: unyt_unit):
return get_openmm_unit_unit(quantity)


## To astropy.units

def quantity_to_astropy_units(quantity: Union[unyt_array, unyt_quantity]):
""" Transform a quantity from unyt to astropy.units."""

from .api_pint import quantity_to_astropy_units as pint_to_astropy_units

pint_quantity = quantity.to_pint()

return pint_to_astropy_units(pint_quantity)


def unit_to_astropy_units(unit: unyt_unit):
""" Transform a unit from unyt to astropy.units."""

from .api_astropy_unit import get_unit as get_astropy_unit

quantity = quantity_to_astropy_units(1.0*unit)

return get_astropy_unit(quantity)


Loading